Download evaluate_validators.py from Offlin33er/doc-fraud-ocr-validate: direct link, hf CLI and curl.
- Browser
- Download file 7.42 kB
-
https://huggingface.co/Offlin33er/doc-fraud-ocr-validate/resolve/main/evaluate_validators.py
- Command line
-
hf download hf://Offlin33er/doc-fraud-ocr-validate/evaluate_validators.py
-
curl -L -o evaluate_validators.py https://huggingface.co/Offlin33er/doc-fraud-ocr-validate/resolve/main/evaluate_validators.py
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() |