TurkishDecisionBenchmark / scripts /validate_dataset.py
muratcanlaloglu
Publish the v0.2 leadboard.
de8702b
Raw History Blame Contribute Delete
8.02 kB
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", "<missing 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())