doc-fraud-ocr-validate / evaluate_validators.py
Offlin33er's picture
Seeded synthetic evaluation harness for the rule layer
f80a3f2 verified
Raw History Blame Contribute Delete
7.42 kB
"""Evaluate the rule layer on seeded synthetic cases with known labels.
Labels are derived at GENERATION time (by construction), never by calling
the validators, so a shared bug cannot hide. Reports per-rule precision,
recall and F1.
Run: python evaluate_validators.py [--n 5000] [--seed 13] [--out results_validators.json]
"""
import argparse
import json
import random
from datetime import date, timedelta
from validators import validate_ssn, validate_date, validate_name, validate_mrz_line2
def gen_valid_ssn(rng):
area = rng.choice([a for a in range(1, 900) if a != 666])
return "%03d-%02d-%04d" % (area, rng.randint(1, 99), rng.randint(1, 9999))
def gen_invalid_ssn(rng):
mode = rng.choice(["zero_area", "area666", "area9xx", "zero_group", "zero_serial", "malformed"])
if mode == "zero_area":
return "000-%02d-%04d" % (rng.randint(1, 99), rng.randint(1, 9999)), "invalid_area"
if mode == "area666":
return "666-%02d-%04d" % (rng.randint(1, 99), rng.randint(1, 9999)), "invalid_area"
if mode == "area9xx":
return "%03d-%02d-%04d" % (rng.randint(900, 999), rng.randint(1, 99), rng.randint(1, 9999)), "invalid_area"
if mode == "zero_group":
return "%03d-00-%04d" % (rng.randint(1, 899), rng.randint(1, 9999)), "invalid_group"
if mode == "zero_serial":
return "%03d-%02d-0000" % (rng.randint(1, 899), rng.randint(1, 99)), "invalid_serial"
style = rng.choice(["too_short", "letters", "no_dashes"])
if style == "too_short":
return "%03d-%02d-%03d" % (rng.randint(1, 899), rng.randint(1, 99), rng.randint(1, 999)), "malformed"
if style == "letters":
return "%03d-AB-%04d" % (rng.randint(1, 899), rng.randint(1, 9999)), "malformed"
return "%03d%02d%04d" % (rng.randint(1, 899), rng.randint(1, 99), rng.randint(1, 9999)), "malformed"
def gen_valid_date(rng, past=True):
end = date(2026, 9, 1) if past else date(2030, 1, 1)
start = date(1940, 1, 1) if past else date(2026, 9, 26)
d = start + timedelta(days=rng.randint(0, max(1, (end - start).days)))
return d.strftime(rng.choice(["%m/%d/%Y", "%Y-%m-%d", "%d.%m.%Y"]))
def gen_invalid_date(rng):
bad = rng.choice(["month13", "day32", "garbage", "future"])
if bad == "month13":
return "13/%02d/%d" % (rng.randint(1, 28), rng.randint(1970, 2020)), "unparseable"
if bad == "day32":
return "%02d/32/%d" % (rng.randint(1, 12), rng.randint(1970, 2020)), "unparseable"
if bad == "garbage":
return rng.choice(["", "N/A", "32/2020", "0101"]), "unparseable"
return "%02d/%02d/%d" % (rng.randint(1, 12), rng.randint(1, 28), rng.randint(2027, 2035)), "future_date"
def gen_valid_name(rng):
first = rng.choice(["JAMES", "MARIA", "CHEN", "OLUWASEUN", "ANNA", "OMAR", "SOFIA"])
last = rng.choice(["SMITH", "OKAFOR", "GARCIA", "O'BRIEN", "NGUYEN", "MUELLER", "KUMARI"])
return first + " " + last
def gen_invalid_name(rng):
return rng.choice(["", " ", "JOHN123", "!!!", "A" * 100, "-9X"]), "bad_characters"
_ALPHABET = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ"
def _compute_check(field):
total = sum((0 if c == "<" else _ALPHABET.index(c)) * (7, 3, 1)[i % 3]
for i, c in enumerate(field))
return str(total % 10)
def gen_valid_mrz(rng):
pno = "".join(rng.choice(_ALPHABET[:36]) for _ in range(rng.randint(6, 9))).ljust(9, "<")
dob = "%02d%02d%02d" % (rng.randint(40, 99), rng.randint(1, 12), rng.randint(1, 28))
exp = "%02d%02d%02d" % (rng.randint(25, 35), rng.randint(1, 12), rng.randint(1, 28))
sex = rng.choice("MF")
per = "<" * 14
pno_ck = _compute_check(pno)
dob_ck = _compute_check(dob)
exp_ck = _compute_check(exp)
per_ck = _compute_check(per)
comp_ck = _compute_check(pno + pno_ck + dob + dob_ck + exp + exp_ck + per + per_ck)
line = pno + pno_ck + "USA" + dob + dob_ck + sex + exp + exp_ck + per + per_ck + comp_ck
assert len(line) == 44, len(line)
return line
def _mrz_checks_ok(line):
return (line[9] == _compute_check(line[0:9])
and line[19] == _compute_check(line[13:19])
and line[27] == _compute_check(line[21:27])
and line[42] == _compute_check(line[28:42])
and line[43] == _compute_check(line[0:10] + line[13:20] + line[21:43]))
def gen_invalid_mrz(rng):
line = list(gen_valid_mrz(rng))
# positions 10-12 are the nationality field, covered by no line-2 check digit
pos = rng.choice([i for i in range(44) if i not in (10, 11, 12)])
old = line[pos]
new = rng.choice([c for c in "0123456789ABC<" if c != old])
line[pos] = new
line = "".join(line)
# a data-position substitution collides with the check digit ~7% of the
# time and would leave the line valid; force the composite digit to
# mismatch so the case is invalid by construction
if _mrz_checks_ok(line):
line = line[:43] + rng.choice([c for c in "0123456789" if c != line[43]])
return line, "char_%d_substituted" % pos
def prf(tp, fp, fn):
precision = tp / (tp + fp) if (tp + fp) else float("nan")
recall = tp / (tp + fn) if (tp + fn) else float("nan")
f1 = (2 * precision * recall / (precision + recall) if (precision + recall) else float("nan"))
return precision, recall, f1
def evaluate_rule(name, cases, validator):
tp = fp = tn = fn = 0
for value, expected_ok, _mode in cases:
ok, _why = validator(value)
if expected_ok and ok:
tp += 1
elif expected_ok and not ok:
fn += 1
elif not expected_ok and ok:
fp += 1
else:
tn += 1
p, r, f = prf(tp, fp, fn)
return {"rule": name, "n": len(cases), "tp": tp, "fp": fp, "tn": tn, "fn": fn,
"precision": round(p, 4), "recall": round(r, 4), "f1": round(f, 4)}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--n", type=int, default=5000, help="cases per rule")
ap.add_argument("--seed", type=int, default=13)
ap.add_argument("--out", default="results_validators.json")
args = ap.parse_args()
rng = random.Random(args.seed)
half = args.n // 2
ssn_cases = [(gen_valid_ssn(rng), True, "valid") for _ in range(half)]
ssn_cases += [(v, False, m) for v, m in (gen_invalid_ssn(rng) for _ in range(args.n - half))]
ref = date(2026, 9, 26)
date_cases = [(gen_valid_date(rng, past=True), True, "valid_past") for _ in range(half)]
date_cases += [(v, False, m) for v, m in (gen_invalid_date(rng) for _ in range(args.n - half))]
name_cases = [(gen_valid_name(rng), True, "valid") for _ in range(half)]
name_cases += [(v, False, m) for v, m in (gen_invalid_name(rng) for _ in range(args.n - half))]
mrz_cases = [(gen_valid_mrz(rng), True, "valid") for _ in range(half)]
mrz_cases += [(v, False, m) for v, m in (gen_invalid_mrz(rng) for _ in range(args.n - half))]
results = [
evaluate_rule("ssn", ssn_cases, validate_ssn),
evaluate_rule("date_past", date_cases,
lambda v: validate_date(v, must_be_past=True, reference=ref)),
evaluate_rule("name", name_cases, validate_name),
evaluate_rule("mrz_line2", mrz_cases, validate_mrz_line2),
]
for r in results:
print(r)
with open(args.out, "w") as f:
json.dump(results, f, indent=2)
print("wrote " + args.out)
if __name__ == "__main__":
main()