"""Build frozen train/temperature/development/test folds from a recipe. Pipeline per recipe part: load (adapter) -> validate -> decontaminate against higher-priority folds -> deduplicate -> cap (sampled by whole family). Folds are processed in priority order test > development > temperature > train, so a state or family used for evaluation can never reach a lower fold. Every loaded row of an evaluation part is claimed, even ones a cap later drops. """ import copy import hashlib import json import multiprocessing import random import re import tomllib from collections import Counter, defaultdict from collections.abc import Sequence from concurrent.futures import ProcessPoolExecutor from datetime import datetime, timezone from pathlib import Path from typing import cast from kev import sources from kev.evaluate import options from kev.types import Example, JSONValue, Label FOLDS = ("test", "development", "temperature", "train") _TOKENIZE: tuple[object, list[str]] | None = None def digest(path: Path) -> str: with path.open("rb") as stream: return hashlib.file_digest(stream, "sha256").hexdigest() def normalized(value: object) -> str: text = value if isinstance(value, str) else json.dumps(value, sort_keys=True, ensure_ascii=False) return re.sub(r"[\W_]+", " ", text.lower()).strip() def state_key(row: Example) -> str: return hashlib.sha256(normalized(row["state"]).encode()).hexdigest() def question_key(row: Example) -> str: question = {key: value for key, value in row["question"].items() if key != "instructions"} signature = normalized(row["state"]) + "\x00" + normalized(row["question"].get("instructions") or "") + "\x00" + normalized(question) return hashlib.sha256(signature.encode()).hexdigest() def family_key(row: Example) -> tuple[str, str]: return str(row["source"]["dataset"]), row["family"] def distribution(row: Example) -> dict[Label, float]: """Target probability per option label (not per position: copies may list options in another order).""" labels = options(row["question"]) target = row["target"] if isinstance(target, list): values = [float(p) for p in target] elif row["question"]["type"] == "noul": values = [1 - float(cast(float, target)), float(cast(float, target))] else: values = [float(type(label) is type(target) and label == target) for label in labels] return dict(zip(labels, values, strict=True)) def merge_duplicates(rows: list[Example]) -> Example: """One row per question; if copies disagree, the target becomes their mean distribution.""" first = rows[0] maps = [distribution(row) for row in rows] if all(mapping == maps[0] for mapping in maps): return first labels = options(first["question"]) mean = [sum(mapping[label] for mapping in maps) / len(maps) for label in labels] result = copy.deepcopy(first) result["label"] = labels[max(range(len(mean)), key=mean.__getitem__)] result["target"] = mean[1] if first["question"]["type"] == "noul" else mean result["source"]["merged_ids"] = [row["id"] for row in rows] return result def sample_families(rows: list[Example], cap: int, seed: int) -> list[Example]: """Keep whole families (all questions of a state) until the cap is reached.""" if len(rows) <= cap: return rows families = sorted({row["family"] for row in rows}) random.Random(seed).shuffle(families) sizes = Counter(row["family"] for row in rows) chosen: set[str] = set() total = 0 for family in families: if total + sizes[family] > cap and total: continue chosen.add(family) total += sizes[family] if total >= cap: break return [row for row in rows if row["family"] in chosen] def expand(patterns: Sequence[str]) -> list[Path]: """Recipe paths may use globs; each must match at least one file.""" paths: list[Path] = [] for pattern in patterns: matches = sorted(Path("/").glob(pattern.lstrip("/"))) if any(c in pattern for c in "*?[") else [Path(pattern)] if not matches or not all(path.is_file() for path in matches): raise FileNotFoundError(f"No input files for {pattern}") paths.extend(matches) return paths def preselect(part: dict[str, object], paths: list[Path], excluded: list[re.Pattern[str]], seed: int) -> set[str]: """Choose groups from (source, group_id) columns only, oversampling each cap by preselect_margin. Final caps are applied after validation, decontamination and deduplication, so the margin absorbs those losses without materializing every state of a multi-million-row source. """ margin = float(cast(float, part["preselect_margin"])) prefix_caps = cast(dict[str, int], part.get("source_caps", {})) per_source = cast(int | None, part.get("per_source_cap")) sizes: Counter[str] = Counter() bucket_of: dict[str, str] = {} for path in paths: for source, group in sources.group_index(path): if any(pattern.search(source) for pattern in excluded): continue prefix = next((prefix for prefix in prefix_caps if source.startswith(prefix)), None) bucket_of[group] = f"prefix:{prefix}" if prefix is not None else f"source:{source}" sizes[group] += 1 buckets: defaultdict[str, list[str]] = defaultdict(list) for group in sorted(bucket_of): buckets[bucket_of[group]].append(group) chosen: set[str] = set() for bucket, members in sorted(buckets.items()): kind, name = bucket.split(":", 1) cap = prefix_caps[name] if kind == "prefix" else per_source if cap is None: chosen.update(members) continue random.Random(f"{seed}:{bucket}").shuffle(members) total = 0 for group in members: if total >= cap * margin: break chosen.add(group) total += sizes[group] return chosen def _init_tokenizer(base_model: str) -> None: global _TOKENIZE from transformers import AutoProcessor from kev.model import answer_codes processor = AutoProcessor.from_pretrained(base_model, local_files_only=True) _TOKENIZE = (processor, answer_codes(processor.tokenizer)[0]) def _count_tokens(rows: Sequence[Example]) -> list[int]: from kev.model import decision_messages assert _TOKENIZE is not None processor, codes = _TOKENIZE texts = [processor.apply_chat_template(decision_messages(row, codes), tokenize=False, # type: ignore[attr-defined] add_generation_prompt=True, enable_thinking=False) for row in rows] encoded = processor.tokenizer(texts, add_special_tokens=False)["input_ids"] # type: ignore[attr-defined] return [len(ids) for ids in encoded] def count_tokens(rows: list[Example], base_model: str, workers: int) -> list[int]: chunks = [rows[start:start + 2000] for start in range(0, len(rows), 2000)] context = multiprocessing.get_context("spawn") # fork after tokenizer threads start can deadlock with ProcessPoolExecutor(workers, mp_context=context, initializer=_init_tokenizer, initargs=(base_model,)) as pool: return [count for counts in pool.map(_count_tokens, chunks) for count in counts] def stats(values: list[int]) -> dict[str, int]: if not values: return {} ordered = sorted(values) return {name: ordered[min(len(ordered) - 1, int(q * len(ordered)))] for name, q in (("p50", 0.5), ("p90", 0.9), ("p99", 0.99), ("max", 1.0))} | {"total": sum(ordered)} def build(recipe_path: Path, output: Path, base_model: str, workers: int, limit: int | None = None) -> dict[str, JSONValue]: from kev.data import validate_row from kev.train import check_partitions recipe = tomllib.loads(recipe_path.read_text()) seed = int(recipe.get("seed", 20260920)) max_tokens = int(recipe.get("max_input_tokens", 8192)) parts = recipe["part"] for part in parts: if part["fold"] not in FOLDS or part["adapter"] not in sources.ADAPTERS: raise ValueError(f"Bad part {part.get('panel')}: fold must be one of {FOLDS}, adapter one of {sources.ADAPTERS}") if output.exists() and any(output.iterdir()): raise FileExistsError(f"Refusing to overwrite a data build: {output}") claimed_states: set[str] = set() claimed_families: set[tuple[str, str]] = set() folds: dict[str, list[Example]] = {fold: [] for fold in FOLDS} report: dict[str, dict[str, JSONValue]] = {} inputs: dict[str, str] = {} for fold in FOLDS: seen_questions: dict[str, list[Example]] = {} seen_ids: set[str] = set() fold_states: set[str] = set() fold_families: set[tuple[str, str]] = set() for index, part in enumerate(p for p in parts if p["fold"] == fold): panel = part["panel"] counts: Counter[str] = Counter() invalid: Counter[str] = Counter() kept: list[Example] = [] excluded = [re.compile(pattern) for pattern in part.get("exclude_sources", [])] paths = expand(part["paths"]) groups = preselect(part, paths, excluded, seed + index) if "preselect_margin" in part else None if groups is not None: counts["preselected_groups"] = len(groups) for path in paths: inputs[str(path)] = digest(path) for loaded, row in enumerate(sources.load(part["adapter"], path, panel, groups)): if limit is not None and loaded >= limit: break counts["loaded"] += 1 if any(pattern.search(row["suite"]) for pattern in excluded): counts["excluded_source"] += 1 continue errors = validate_row(row) if errors: invalid[errors[0]] += 1 continue state, family = state_key(row), family_key(row) if fold != "train": fold_states.add(state) fold_families.add(family) if state in claimed_states or family in claimed_families: counts["dropped_overlaps_higher_fold"] += 1 continue key = question_key(row) if row["id"] in seen_ids: # Same source item in another format/version: keep it under a distinct id. row["id"] = f"{row['id']}~{key[:12]}" counts["renamed_id_collisions"] += 1 if row["id"] in seen_ids: counts["dropped_duplicate"] += 1 continue if fold == "train" and key in seen_questions: if seen_questions[key][0]["source"]["panel"] != panel: counts["dropped_duplicate_of_earlier_part"] += 1 else: seen_questions[key].append(row) continue seen_questions.setdefault(key, []).append(row) seen_ids.add(row["id"]) kept.append(row) if fold == "train": # Collapse identical questions; conflicting labels become one soft target (their label distribution). merged: list[Example] = [] for row in kept: group = seen_questions[question_key(row)] counts["dropped_duplicate"] += len(group) - 1 result = merge_duplicates(group) if result is not row: counts["merged_label_conflicts"] += 1 merged.append(result) kept = merged else: # Evaluation suites stay as published: repeated inputs with different labels are # deliberate (e.g. kev "unknowable" pairs test calibrated 50/50 answers). counts["kept_repeated_questions"] = sum(len(group) - 1 for group in seen_questions.values() if group[0]["source"]["panel"] == panel) counts["valid_unique"] = len(kept) for suite_prefix, cap in part.get("source_caps", {}).items(): matching = [row for row in kept if row["suite"].startswith(suite_prefix)] sampled = {id(row) for row in sample_families(matching, int(cap), seed + index)} kept = [row for row in kept if not row["suite"].startswith(suite_prefix) or id(row) in sampled] counts[f"after_source_cap:{suite_prefix}"] = len(kept) if "per_source_cap" in part: prefixes = tuple(part.get("source_caps", {})) by_source: defaultdict[str, list[Example]] = defaultdict(list) for row in kept: by_source[row["suite"]].append(row) kept = [row for suite, rows in sorted(by_source.items()) for row in (rows if suite.startswith(prefixes) and prefixes else sample_families(rows, int(part["per_source_cap"]), seed + index))] counts["after_per_source_cap"] = len(kept) counts["sources"] = len(by_source) if "cap" in part: kept = sample_families(kept, int(part["cap"]), seed + index) counts["selected"] = len(kept) folds[fold].extend(kept) report[f"{fold}/{panel}"] = {"adapter": part["adapter"], "paths": part["paths"], **counts, "invalid": cast(JSONValue, dict(invalid))} # Evaluation folds claim every loaded state/family, not only the sampled ones. claimed_states |= fold_states | {state_key(row) for row in folds[fold]} claimed_families |= fold_families | {family_key(row) for row in folds[fold]} everything = [row for fold in FOLDS for row in folds[fold]] lengths = count_tokens(everything, base_model, workers) too_long: Counter[str] = Counter() for row, length in zip(everything, lengths, strict=True): row["source"]["input_tokens"] = length for fold in FOLDS: before = len(folds[fold]) folds[fold] = [row for row in folds[fold] if cast(int, row["source"]["input_tokens"]) <= max_tokens] too_long[fold] = before - len(folds[fold]) check_partitions(folds) output.mkdir(parents=True, exist_ok=True) files: dict[str, JSONValue] = {} summary: dict[str, JSONValue] = {} for fold in FOLDS: path = output / f"{fold}.jsonl" path.write_text("".join(json.dumps(row, ensure_ascii=False) + "\n" for row in folds[fold])) files[path.name] = digest(path) panels: defaultdict[str, Counter[str]] = defaultdict(Counter) for row in folds[fold]: panels[str(row["source"]["panel"])][row["question"]["type"]] += 1 summary[fold] = { "questions": len(folds[fold]), "states": len({state_key(row) for row in folds[fold]}), "soft_targets": sum(isinstance(row["target"], list) for row in folds[fold]), "dropped_over_max_tokens": too_long[fold], "input_tokens": cast(JSONValue, stats([cast(int, row["source"]["input_tokens"]) for row in folds[fold]])), "panels": {panel: dict(types) for panel, types in sorted(panels.items())}, } package = Path(__file__).resolve().parent manifest: dict[str, JSONValue] = { "name": recipe.get("name", output.name), "created": datetime.now(timezone.utc).isoformat(), "recipe": str(recipe_path.resolve()), "recipe_sha256": digest(recipe_path), "seed": seed, "max_input_tokens": max_tokens, "limit_per_file": limit, "base_model": base_model, "summary": summary, "parts": cast(JSONValue, report), "inputs_sha256": cast(JSONValue, inputs), "outputs_sha256": files, "code_sha256": {name: digest(package / name) for name in ("build.py", "sources.py", "data.py", "model.py")}, } (output / "recipe.toml").write_text(recipe_path.read_text()) (output / "manifest.json").write_text(json.dumps(manifest, indent=2, ensure_ascii=False) + "\n") return manifest