import argparse import sys from collections import Counter, defaultdict from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from benchmark.dataset import ( # noqa: E402 BENCHMARK_VERSION, VERSIONS, load_cases, load_tasks, sha256_file, ) DIFFICULTIES = {"easy", "medium", "hard"} BASE_FIELDS = { "id", "version", "split", "category", "difficulty", "task_id", "state", "expected", "valid_answers", "scored", } V02_FIELDS = {"domain", "phenomena", "group_id", "answer_position", "review_status", "annotations"} V02_CATEGORIES = { "negation", "correction", "temporal_reasoning", "distractor", "implicit_intent", "coreference_scope", "conditional", "reported_speech", "noisy_turkish", "numeric_date", "ambiguous", } V02_DOMAINS = { "subscription", "ecommerce", "banking", "telecom", "shipping", "health", "public_services", "travel", } POSITIONS = {"start", "middle", "end", "na"} REVIEW_STATUSES = {"draft", "reviewed", "adjudicated"} MAX_OPTIONS = 20 MAX_IMBALANCE = 2.0 def validate_tasks(tasks): errors = [] for name, task in tasks.items(): criteria = task.get("criteria") or {} if task.get("type") != "choice" or len(criteria) < 2: errors.append(f"task {name}: must be a 'choice' task with >= 2 criteria") if len(criteria) > MAX_OPTIONS: errors.append(f"task {name}: {len(criteria)} options > {MAX_OPTIONS}") if not str(task.get("instructions", "")).strip(): errors.append(f"task {name}: empty instructions") return errors def validate_case(r, tasks, version, split): rid = r.get("id", "") required = BASE_FIELDS | V02_FIELDS missing = required - r.keys() if missing: return [f"{rid}: missing fields {sorted(missing)}"] errors = [] if r["version"] != version: errors.append(f"{rid}: version {r['version']!r} != {version!r}") if r["split"] != split: errors.append(f"{rid}: split {r['split']!r} in {split} file") if r["difficulty"] not in DIFFICULTIES: errors.append(f"{rid}: unknown difficulty {r['difficulty']!r}") if not str(r["state"]).strip(): errors.append(f"{rid}: empty state") if r["task_id"] not in tasks: return errors + [f"{rid}: unknown task_id {r['task_id']!r}"] criteria = tasks[r["task_id"]]["criteria"] for answer in r["valid_answers"]: if answer not in criteria: errors.append(f"{rid}: valid answer {answer!r} not in task criteria") if r["scored"]: if r["expected"] not in criteria: errors.append(f"{rid}: expected {r['expected']!r} not in task criteria") if r["valid_answers"] != [r["expected"]]: errors.append(f"{rid}: scored case must have valid_answers == [expected]") else: if r["expected"] is not None: errors.append(f"{rid}: unscored case must have expected == null") if len(r["valid_answers"]) < 2: errors.append(f"{rid}: unscored case needs >= 2 valid answers") if len(r["valid_answers"]) >= len(criteria): errors.append(f"{rid}: every class is valid, the case measures nothing") if r["category"] not in V02_CATEGORIES: errors.append(f"{rid}: unknown category {r['category']!r}") if r["domain"] not in V02_DOMAINS: errors.append(f"{rid}: unknown domain {r['domain']!r}") prefix = r["task_id"].split(".", 1)[0] if prefix not in ("common", r["domain"]): errors.append(f"{rid}: task {r['task_id']!r} does not belong to domain {r['domain']!r}") if r["answer_position"] not in POSITIONS: errors.append(f"{rid}: answer_position must be one of {sorted(POSITIONS)}") if r["review_status"] not in REVIEW_STATUSES: errors.append(f"{rid}: review_status must be one of {sorted(REVIEW_STATUSES)}") if not isinstance(r["phenomena"], list): errors.append(f"{rid}: phenomena must be a list") ann = r["annotations"] or {} for key, label in ann.items(): if label is not None and label not in criteria: errors.append(f"{rid}: annotation {key}={label!r} not in task criteria") if r["review_status"] == "adjudicated": if r["scored"] and ann.get("adjudicated") != r["expected"]: errors.append(f"{rid}: adjudicated label differs from expected") return errors def validate_groups(rows): errors = [] groups = defaultdict(list) for r in rows: if r.get("group_id"): groups[r["group_id"]].append(r) for gid, members in groups.items(): if len(members) < 2: errors.append(f"group {gid}: needs >= 2 cases") if len({m["task_id"] for m in members}) > 1: errors.append(f"group {gid}: mixes task_ids") labels = {m["expected"] for m in members if m["scored"]} if len(labels) < 2: errors.append(f"group {gid}: contrast group must have >= 2 distinct labels") return errors, groups def label_warnings(rows, tasks): warnings = [] for task_id in tasks: counts = Counter(r["expected"] for r in rows if r["scored"] and r["task_id"] == task_id) if not counts: continue missing = [c for c in tasks[task_id]["criteria"] if c not in counts] if missing: warnings.append(f"{task_id}: no scored case for {missing}") elif max(counts.values()) > MAX_IMBALANCE * min(counts.values()): warnings.append(f"{task_id}: label imbalance {dict(counts)}") return warnings def main(): parser = argparse.ArgumentParser() parser.add_argument("--version", default=BENCHMARK_VERSION, choices=sorted(VERSIONS)) parser.add_argument("--dataset", type=Path, help="validate this file as the public split only") args = parser.parse_args() spec = VERSIONS[args.version] tasks = load_tasks(spec["tasks"]) errors = validate_tasks(tasks) splits = {"public": args.dataset} if args.dataset else spec["splits"] rows = [] present = {} for split, path in splits.items(): if not path.exists(): continue split_rows = load_cases(path) present[split] = (path, len(split_rows)) for r in split_rows: errors += validate_case(r, tasks, args.version, split) rows += split_rows for dup, n in Counter(r.get("id") for r in rows).items(): if n > 1: errors.append(f"duplicate id {dup} ({n}x)") for dup, n in Counter(r.get("state") for r in rows).items(): if n > 1: errors.append(f"duplicate state ({n}x): {dup}") group_errors, groups = validate_groups(rows) errors += group_errors if errors: for e in errors: print(f"ERROR: {e}") print(f"\n{len(errors)} error(s).") return 1 scored = [r for r in rows if r["scored"]] print(f"OK v{args.version}: {len(rows)} total cases ({len(scored)} scored)") for split, (path, n) in present.items(): print(f" {split:7} {n:4} sha256 {sha256_file(path)}") print(f" tasks sha256 {sha256_file(spec['tasks'])}") print("Categories:", dict(Counter(r["category"] for r in rows))) print("Difficulty (scored):", dict(Counter(r["difficulty"] for r in scored))) print("Domains:", dict(Counter(r["domain"] for r in rows))) print("Answer position:", dict(Counter(r["answer_position"] for r in scored))) print("Review status:", dict(Counter(r["review_status"] for r in rows))) grouped = sum(len(m) for m in groups.values()) print(f"Contrast groups: {len(groups)} covering {grouped}/{len(rows)} cases") print("Labels (scored):") for task_id in tasks: counts = Counter(r["expected"] for r in scored if r["task_id"] == task_id) if counts: print(f" {task_id:26} {dict(counts)}") for w in label_warnings(rows, tasks): print(f"WARNING: {w}") return 0 if __name__ == "__main__": sys.exit(main())