"""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))