From 9658faa4beb5d11444f25772de541028821221c8 Mon Sep 17 00:00:00 2001 From: Daniel Maddern Date: Fri, 14 Aug 2026 01:12:12 +0700 Subject: [PATCH] Add full VAE decode comparator --- tools/compare_vae_full_decode.py | 58 ++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) create mode 100644 tools/compare_vae_full_decode.py diff --git a/tools/compare_vae_full_decode.py b/tools/compare_vae_full_decode.py new file mode 100644 index 0000000..936bc50 --- /dev/null +++ b/tools/compare_vae_full_decode.py @@ -0,0 +1,58 @@ +"""Compare direct and upstream full VAE decode on a saved latent.""" + +import argparse +import gc +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("--frames", type=int) +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 pixelize(decoded: torch.Tensor, frames: int) -> torch.Tensor: + return ((decoded[:, :, :frames].clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu() + + +def summarize(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({ + "shape": tuple(a.shape), + "max": float(diff.max()), + "mean": float(diff.mean()), + "top_max": sorted(frame_rows, key=lambda row: row[1], reverse=True)[:12], + "top_mean": sorted(frame_rows, key=lambda row: row[2], reverse=True)[:12], + }, 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") +frames = args.frames or int(state.get("frames", latent.shape[2] * 4)) + +with torch.inference_mode(): + direct = DirectVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda").eval() + direct_pixels = pixelize(direct.decode(latent.to(next(direct.parameters()).dtype)), frames) + del direct + torch.cuda.empty_cache() + gc.collect() + + upstream = UpstreamVAE().to("cuda").eval() + upstream.load_state_dict(load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda"), strict=True) + upstream_pixels = pixelize(upstream.decode(latent.to(next(upstream.parameters()).dtype)), frames) + +summarize(direct_pixels, upstream_pixels)