mad-bot's picture
Upload folder using huggingface_hub (part 2)
eae424a verified
Raw
History Blame Contribute Delete
6.38 kB
#!/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()