en-or-not / scripts /evaluate_onnx.py
Mitchins's picture
Publish en-or-not OOD benchmark
3a559ac
Raw
History Blame Contribute Delete
3.41 kB
#!/usr/bin/env python3
"""Evaluate model.onnx against a JSONL set with text and label fields."""
import argparse
import json
import math
from collections import defaultdict
import numpy as np
import onnxruntime as ort
def predict(session, text):
encoded = text.encode("utf-8", errors="ignore")[:256]
values = np.zeros((1, 256), dtype=np.int64)
values[0, : len(encoded)] = list(encoded)
logits = session.run(["logits"], {"input_bytes": values})[0][0]
return "EN" if int(np.argmax(logits)) == 1 else "NOT-EN", float(logits[1] - logits[0])
def main():
parser = argparse.ArgumentParser()
parser.add_argument("dataset", help="JSONL file containing text and label")
parser.add_argument("--model", default="model.onnx")
parser.add_argument("--max-errors", type=int, default=20)
args = parser.parse_args()
session = ort.InferenceSession(args.model, providers=["CPUExecutionProvider"])
rows = [json.loads(line) for line in open(args.dataset, encoding="utf-8") if line.strip()]
correct = 0
by_category = defaultdict(lambda: [0, 0])
by_language = defaultdict(lambda: [0, 0])
errors = []
predictions = []
for row in rows:
prediction, margin = predict(session, row["text"])
predictions.append(prediction)
category = row.get("category", "uncategorized")
language = row.get("language", "uncategorized")
by_category[category][1] += 1
by_language[language][1] += 1
correct += prediction == row["label"]
by_category[category][0] += prediction == row["label"]
by_language[language][0] += prediction == row["label"]
if prediction != row["label"]:
errors.append({"id": row["id"], "text": row["text"], "expected": row["label"], "predicted": prediction, "margin": round(margin, 4)})
accuracy = correct / len(rows)
margin = 1.96 * math.sqrt(accuracy * (1 - accuracy) / len(rows))
print(f"accuracy: {correct}/{len(rows)} = {accuracy:.4%} (95% Wald CI {max(0, accuracy - margin):.4%}-{min(1, accuracy + margin):.4%})")
labels = {label: [0, 0] for label in ("EN", "NOT-EN")}
for row, prediction in zip(rows, predictions):
labels[row["label"]][0] += prediction == row["label"]
labels[row["label"]][1] += 1
print("class accuracy:")
for label in labels:
hit, total = labels[label]
print(f" {label}: {hit}/{total} = {hit / total:.4%}")
macro = sum(hit / total for hit, total in by_language.values()) / len(by_language)
non_english_languages = [lang for lang in by_language if lang != "en"]
non_english_macro = sum(by_language[lang][0] / by_language[lang][1] for lang in non_english_languages) / len(non_english_languages)
print(f"language macro accuracy: {macro:.4%} across {len(by_language)} languages")
print(f"non-English language macro accuracy: {non_english_macro:.4%} across {len(non_english_languages)} languages")
print("category accuracy:")
for category in sorted(by_category):
hit, total = by_category[category]
print(f" {category}: {hit}/{total} = {hit / total:.4%}")
print(f"errors: {len(errors)}")
for error in errors[: args.max_errors]:
print(json.dumps(error, ensure_ascii=False))
if len(errors) > args.max_errors:
print(f"... {len(errors) - args.max_errors} additional errors omitted")
if __name__ == "__main__":
main()