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