32 lines
1.1 KiB
Python
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))
|