matilda-jev-fp4 / kev /build.py
yue-maincode's picture
Upload validated MATILDA JEV FP4 model and Decision Index scores
c69aaec verified
Raw History Blame Contribute Delete
16.5 kB
"""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