Spaces:
Running on Zero
Running on Zero
Download scripts/validate_dataset.py from muratcanlaloglu/TurkishDecisionBenchmark: direct link, hf CLI and curl.
- Browser
- Download file 8.02 kB
-
https://huggingface.co/spaces/muratcanlaloglu/TurkishDecisionBenchmark/resolve/main/scripts/validate_dataset.py
- Command line
-
hf download hf://spaces/muratcanlaloglu/TurkishDecisionBenchmark/scripts/validate_dataset.py
-
curl -L -o validate_dataset.py https://huggingface.co/spaces/muratcanlaloglu/TurkishDecisionBenchmark/resolve/main/scripts/validate_dataset.py
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()) | |