#!/usr/bin/env python3 """Compare two raw little-endian f32 tensors and emit an auditable record.""" from __future__ import annotations import argparse import array import hashlib import json import math import sys from pathlib import Path def read_tensor(path: Path) -> array.array[float]: data = path.read_bytes() if len(data) % 4: raise SystemExit(f"{path} is not a whole number of f32 values") values = array.array("f") values.frombytes(data) if sys.byteorder != "little": values.byteswap() return values def sha256(path: Path) -> str: return hashlib.sha256(path.read_bytes()).hexdigest() def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("reference", type=Path) parser.add_argument("candidate", type=Path) parser.add_argument("--output", required=True, type=Path) parser.add_argument("--max-relative-l2", type=float) parser.add_argument("--min-cosine", type=float) parser.add_argument("--row-size", type=int) parser.add_argument("--min-row-cosine", type=float) args = parser.parse_args() if args.min_row_cosine is not None and args.row_size is None: raise SystemExit("--min-row-cosine requires --row-size") reference = read_tensor(args.reference) candidate = read_tensor(args.candidate) if len(reference) != len(candidate): raise SystemExit( f"length mismatch: {args.reference} has {len(reference)}, " f"{args.candidate} has {len(candidate)}" ) squared_error = 0.0 squared_reference = 0.0 squared_candidate = 0.0 dot = 0.0 absolute_error = 0.0 max_absolute_error = -1.0 max_absolute_index = 0 for index, (expected, actual) in enumerate(zip(reference, candidate)): difference = float(actual) - float(expected) squared_error += difference * difference squared_reference += float(expected) * float(expected) squared_candidate += float(actual) * float(actual) dot += float(expected) * float(actual) absolute_error += abs(difference) if abs(difference) > max_absolute_error: max_absolute_error = abs(difference) max_absolute_index = index relative_l2 = math.sqrt(squared_error) / max(math.sqrt(squared_reference), sys.float_info.min) cosine = dot / max( math.sqrt(squared_reference) * math.sqrt(squared_candidate), sys.float_info.min, ) row_metrics = None if args.row_size is not None: if args.row_size <= 0 or len(reference) % args.row_size: raise SystemExit("--row-size must be positive and divide the tensor length") minimum_cosine = math.inf minimum_cosine_row = 0 maximum_relative_l2 = -math.inf maximum_relative_l2_row = 0 for row in range(len(reference) // args.row_size): start = row * args.row_size end = start + args.row_size row_reference = reference[start:end] row_candidate = candidate[start:end] row_dot = sum(float(a) * float(b) for a, b in zip(row_reference, row_candidate)) row_reference_squared = sum(float(value) ** 2 for value in row_reference) row_candidate_squared = sum(float(value) ** 2 for value in row_candidate) row_error_squared = sum( (float(actual) - float(expected)) ** 2 for expected, actual in zip(row_reference, row_candidate) ) row_cosine = row_dot / max( math.sqrt(row_reference_squared) * math.sqrt(row_candidate_squared), sys.float_info.min, ) row_relative_l2 = math.sqrt(row_error_squared) / max( math.sqrt(row_reference_squared), sys.float_info.min ) if row_cosine < minimum_cosine: minimum_cosine = row_cosine minimum_cosine_row = row if row_relative_l2 > maximum_relative_l2: maximum_relative_l2 = row_relative_l2 maximum_relative_l2_row = row row_metrics = { "row_size": args.row_size, "rows": len(reference) // args.row_size, "minimum_cosine": minimum_cosine, "minimum_cosine_row": minimum_cosine_row, "maximum_relative_l2": maximum_relative_l2, "maximum_relative_l2_row": maximum_relative_l2_row, } checks = { "max_relative_l2": ( None if args.max_relative_l2 is None else relative_l2 <= args.max_relative_l2 ), "min_cosine": ( None if args.min_cosine is None else cosine >= args.min_cosine ), "min_row_cosine": ( None if args.min_row_cosine is None else row_metrics is not None and row_metrics["minimum_cosine"] >= args.min_row_cosine ), } passed = all(value is not False for value in checks.values()) record = { "schema_version": 1, "reference": str(args.reference).replace("\\", "/"), "reference_sha256": sha256(args.reference), "candidate": str(args.candidate).replace("\\", "/"), "candidate_sha256": sha256(args.candidate), "elements": len(reference), "relative_l2": relative_l2, "cosine": cosine, "mean_absolute_error": absolute_error / len(reference), "max_absolute_error": max_absolute_error, "max_absolute_index": max_absolute_index, "thresholds": { "max_relative_l2": args.max_relative_l2, "min_cosine": args.min_cosine, "min_row_cosine": args.min_row_cosine, }, "row_metrics": row_metrics, "checks": checks, "passed": passed, } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(record, indent=2) + "\n", encoding="utf-8") row_text = ( "" if row_metrics is None else f", minimum row cosine {row_metrics['minimum_cosine']:.9f}" ) print( f"relative L2 {relative_l2:.6g}, cosine {cosine:.9f}{row_text}, " f"max abs {max_absolute_error:.6g} at {max_absolute_index}: " f"{'PASS' if passed else 'FAIL'}" ) if not passed: raise SystemExit(1) if __name__ == "__main__": main()