h3-blackwell-runtime/tools/compare_generation_latents.py
2026-08-20 17:39:44 +07:00

32 lines
1.1 KiB
Python

"""Compare saved hot-runtime video and audio latents against one reference run."""
import argparse
import json
import torch
parser = argparse.ArgumentParser()
parser.add_argument("reference")
parser.add_argument("candidates", nargs="+")
args = parser.parse_args()
reference = torch.load(args.reference, map_location="cpu", weights_only=False)
results = {}
for path in args.candidates:
candidate = torch.load(path, map_location="cpu", weights_only=False)
metrics = {}
for name in ("latent", "audio_latent"):
expected = reference[name].float()
actual = candidate[name].float()
delta = actual - expected
metrics[name] = {
"max_abs": delta.abs().max().item(),
"mean_abs": delta.abs().mean().item(),
"rmse": delta.square().mean().sqrt().item(),
"relative_rmse": (delta.square().mean().sqrt() / expected.square().mean().sqrt()).item(),
"cosine": torch.nn.functional.cosine_similarity(actual.flatten(), expected.flatten(), dim=0).item(),
}
results[path] = metrics
print(json.dumps({"reference": args.reference, "results": results}, indent=2))