File size: 4,976 Bytes
c69aaec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""Build and validate decision JSONL folds.

`kev-data build RECIPE OUTPUT` builds frozen folds from the archived sources
(see kev.build and configs/data/*.toml).

`kev-data validate train.jsonl dev.jsonl ...` checks every row against the
Example schema, summarizes suites/question types, and rejects ID or
(dataset, family) overlap between the given files.
"""

import argparse
import json
import math
import os
from collections import Counter
from pathlib import Path

from kev.evaluate import options, read_rows
from kev.model import BASE_MODEL, MAX_OPTIONS
from kev.types import Example


def validate_row(row: Example) -> list[str]:
    errors: list[str] = []
    for key in ("id", "suite", "family"):
        if not isinstance(row.get(key), str) or not row.get(key):
            errors.append(f"missing {key}")
    if not isinstance(row.get("source"), dict):
        errors.append("missing source")
    question = row.get("question")
    if not isinstance(question, dict) or question.get("type") not in ("choice", "noul", "score"):
        return errors + ["question.type must be choice, noul or score"]
    if question["type"] == "choice" and (not isinstance(question.get("criteria"), dict) or not question["criteria"]):
        return errors + ["choice questions need a nonempty criteria object"]
    if question["type"] == "score" and (not isinstance(question.get("criteria"), list) or len(question["criteria"]) < 2):
        return errors + ["score questions need at least two criteria levels"]
    labels = options(question)
    if len(labels) > MAX_OPTIONS:
        errors.append(f"more than {MAX_OPTIONS} options")
    target, label = row.get("target"), row.get("label")
    if isinstance(target, list):
        if len(target) != len(labels) or any(not isinstance(p, (int, float)) or not math.isfinite(p) or p < 0 for p in target) \
                or abs(sum(target) - 1) > 1e-6:
            errors.append("soft target must be a distribution over the options")
    elif question["type"] == "noul":
        if isinstance(target, bool):
            pass
        elif not isinstance(target, (int, float)) or not 0 <= float(target) <= 1:
            errors.append("noul target must be a bool or probability")
    elif not any(type(target) is type(value) and target == value for value in labels):
        errors.append("target is not an option")
    if not any(type(label) is type(value) and label == value for value in labels):
        errors.append("hard label is not an option")
    return errors


def validate(paths: list[Path]) -> bool:
    from kev.train import check_partitions

    ok = True
    partitions: dict[str, list[Example]] = {}
    for path in paths:
        rows = read_rows(path)
        partitions[str(path)] = rows
        failures = Counter[str]()
        for row in rows:
            for error in validate_row(row):
                failures[error] += 1
        summary = {
            "rows": len(rows),
            "suites": dict(Counter(row.get("suite") for row in rows).most_common(15)),
            "types": dict(Counter(row["question"]["type"] for row in rows if isinstance(row.get("question"), dict))),
            "images": sum(bool(row.get("images")) for row in rows),
            "errors": dict(failures),
        }
        print(json.dumps({str(path): summary}, indent=2, ensure_ascii=False))
        ok = ok and not failures
    if len(partitions) > 1:
        try:
            check_partitions(partitions)
            print("No ID or family overlap between files.")
        except ValueError as error:
            print(f"Overlap: {error}")
            ok = False
    return ok


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    commands = parser.add_subparsers(dest="command", required=True)
    check = commands.add_parser("validate", help="Schema, summary and overlap checks for JSONL folds")
    check.add_argument("paths", type=Path, nargs="+")
    make = commands.add_parser("build", help="Build frozen folds from a recipe TOML")
    make.add_argument("recipe", type=Path)
    make.add_argument("output", type=Path, help="New directory for the build")
    make.add_argument("--base-model", default=BASE_MODEL, help="Tokenizer used to count prompt tokens")
    make.add_argument("--workers", type=int, default=min(16, os.cpu_count() or 1))
    make.add_argument("--limit", type=int, help="Read at most this many questions per file (quick pipeline check)")
    args = parser.parse_args()
    if args.command == "validate" and not validate(args.paths):
        raise SystemExit(1)
    if args.command == "build":
        from kev.build import build

        manifest = build(args.recipe, args.output, args.base_model, args.workers, args.limit)
        print(json.dumps({"summary": manifest["summary"], "parts": manifest["parts"]}, indent=2, ensure_ascii=False))
        print(f"Wrote {args.output}")


if __name__ == "__main__":
    main()