From 6dd0b0e737a87a9fcfb0f04eeb4d2b04f407d0b3 Mon Sep 17 00:00:00 2001 From: Daniel Maddern Date: Fri, 14 Aug 2026 01:20:27 +0700 Subject: [PATCH] Add VAE tiled clip comparator --- tools/compare_vae_tiled_clip.py | 51 +++++++++++++++++++++++++++++++++ 1 file changed, 51 insertions(+) create mode 100644 tools/compare_vae_tiled_clip.py diff --git a/tools/compare_vae_tiled_clip.py b/tools/compare_vae_tiled_clip.py new file mode 100644 index 0000000..76034c4 --- /dev/null +++ b/tools/compare_vae_tiled_clip.py @@ -0,0 +1,51 @@ +"""Compare direct/upstream VAE tiled decode for one temporal latent clip.""" + +import argparse +import sys +from pathlib import Path + +import torch +from safetensors.torch import load_file + +from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE as DirectVAE + + +parser = argparse.ArgumentParser() +parser.add_argument("--latent", type=Path, required=True) +parser.add_argument("--start", type=int, default=0) +parser.add_argument("--tokens", type=int, default=8) +parser.add_argument("--comfy-path", default="/opt/ComfyUI") +args = parser.parse_args() + +sys.path.insert(0, args.comfy_path) +from h3_blackwell_runtime.upstream_vae import MiniMaxH3VideoVAE as UpstreamVAE # noqa: E402 + + +def stats(name: str, a: torch.Tensor, b: torch.Tensor) -> None: + diff = (a.float() - b.float()).abs() + frame_rows = [] + for index in range(diff.shape[2]): + frame = diff[:, :, index] + frame_rows.append((index, float(frame.max()), float(frame.mean()), float(torch.quantile(frame.flatten(), 0.99)))) + print({ + "stage": name, + "shape": tuple(a.shape), + "max": float(diff.max()), + "mean": float(diff.mean()), + "top_mean": sorted(frame_rows, key=lambda row: row[2], reverse=True)[:8], + }, flush=True) + + +state = torch.load(args.latent, map_location="cuda", weights_only=False) +latent = state["latent"].to("cuda") if isinstance(state, dict) else state.to("cuda") + +direct = DirectVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda", tiling=True).eval() +upstream = UpstreamVAE(tiling=True).to("cuda").eval() +upstream.load_state_dict(load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda"), strict=True) + +with torch.inference_mode(): + z_d = latent[:, :, args.start:args.start + args.tokens].to(next(direct.parameters()).dtype) + z_u = latent[:, :, args.start:args.start + args.tokens].to(next(upstream.parameters()).dtype) + z_d = z_d * direct.latents_std.view(1, -1, 1, 1, 1).to(z_d) + direct.latents_mean.view(1, -1, 1, 1, 1).to(z_d) + z_u = z_u * upstream.latents_std.view(1, -1, 1, 1, 1).to(z_u) + upstream.latents_mean.view(1, -1, 1, 1, 1).to(z_u) + stats("tiled_clip", direct.tiled_decode(z_d), upstream.tiled_decode(z_u))