517 lines
18 KiB
Python
517 lines
18 KiB
Python
"""Shared data contracts for the purpose-classifier pipeline."""
|
|||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import hashlib
|
||
|
|
import json
|
||
|
|
import math
|
||
|
|
import re
|
||
|
|
import unicodedata
|
||
|
|
from collections import Counter, defaultdict
|
||
|
|
from dataclasses import dataclass
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any, Callable, Iterable, Sequence
|
||
|
|
|
||
|
|
|
||
|
|
LABELS = (
|
||
|
|
"planning",
|
||
|
|
"backendImpl",
|
||
|
|
"frontendImpl",
|
||
|
|
"quickFix",
|
||
|
|
"refactor",
|
||
|
|
"debugging",
|
||
|
|
"review",
|
||
|
|
"writing",
|
||
|
|
)
|
||
|
|
LABEL_SET = frozenset(LABELS)
|
||
|
|
SOURCE_FIELDS = frozenset(
|
||
|
|
{"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"}
|
||
|
|
)
|
||
|
|
SLICES = frozenset(
|
||
|
|
{"core", "boundary", "mixed", "pasted-context", "vague-eval"}
|
||
|
|
)
|
||
|
|
HARD_SLICES = frozenset({"boundary", "mixed", "pasted-context", "vague-eval"})
|
||
|
|
WORD_RE = re.compile(r"\w+", re.UNICODE)
|
||
|
|
|
||
|
|
|
||
|
|
class DataError(ValueError):
|
||
|
|
"""A deterministic data-contract failure."""
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True)
|
||
|
|
class SourceRecord:
|
||
|
|
value: dict[str, Any]
|
||
|
|
source: Path
|
||
|
|
line: int
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True)
|
||
|
|
class Duplicate:
|
||
|
|
dropped: SourceRecord
|
||
|
|
matched_prompt_hash: str
|
||
|
|
kind: str
|
||
|
|
similarity: float
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True)
|
||
|
|
class CurationResult:
|
||
|
|
records: list[SourceRecord]
|
||
|
|
duplicates: list[Duplicate]
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True)
|
||
|
|
class SplitResult:
|
||
|
|
train: list[SourceRecord]
|
||
|
|
validation: list[SourceRecord]
|
||
|
|
test: list[SourceRecord]
|
||
|
|
fixture_count: int
|
||
|
|
|
||
|
|
@property
|
||
|
|
def logical_test_count(self) -> int:
|
||
|
|
return len(self.test) + self.fixture_count
|
||
|
|
|
||
|
|
|
||
|
|
def normalize_prompt(prompt: str) -> str:
|
||
|
|
"""Match the runtime's whitespace collapse and add stable Unicode normalization."""
|
||
|
|
|
||
|
|
normalized = unicodedata.normalize("NFKC", prompt)
|
||
|
|
return " ".join(normalized.split())
|
||
|
|
|
||
|
|
|
||
|
|
def normalized_key(prompt: str) -> str:
|
||
|
|
return normalize_prompt(prompt).casefold()
|
||
|
|
|
||
|
|
|
||
|
|
def prompt_hash(prompt: str) -> str:
|
||
|
|
return hashlib.sha256(normalized_key(prompt).encode("utf-8")).hexdigest()
|
||
|
|
|
||
|
|
|
||
|
|
def file_sha256(path: Path) -> str:
|
||
|
|
digest = hashlib.sha256()
|
||
|
|
with path.open("rb") as handle:
|
||
|
|
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
|
||
|
|
digest.update(chunk)
|
||
|
|
return digest.hexdigest()
|
||
|
|
|
||
|
|
|
||
|
|
def canonical_json(record: dict[str, Any]) -> str:
|
||
|
|
return json.dumps(record, ensure_ascii=False, separators=(",", ":"))
|
||
|
|
|
||
|
|
|
||
|
|
def jsonl_bytes(records: Iterable[dict[str, Any]]) -> bytes:
|
||
|
|
return ("".join(f"{canonical_json(record)}\n" for record in records)).encode("utf-8")
|
||
|
|
|
||
|
|
|
||
|
|
def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None:
|
||
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||
|
|
path.write_bytes(jsonl_bytes(records))
|
||
|
|
|
||
|
|
|
||
|
|
def write_json(path: Path, value: Any) -> None:
|
||
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||
|
|
path.write_text(
|
||
|
|
json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||
|
|
encoding="utf-8",
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def load_jsonl(path: Path) -> list[dict[str, Any]]:
|
||
|
|
records: list[dict[str, Any]] = []
|
||
|
|
try:
|
||
|
|
lines = path.read_text(encoding="utf-8").splitlines()
|
||
|
|
except (OSError, UnicodeError) as exc:
|
||
|
|
raise DataError(f"{path}: cannot read UTF-8 JSONL: {exc}") from exc
|
||
|
|
|
||
|
|
for line_number, line in enumerate(lines, 1):
|
||
|
|
if not line.strip():
|
||
|
|
raise DataError(f"{path}:{line_number}: blank JSONL line")
|
||
|
|
try:
|
||
|
|
value = json.loads(line)
|
||
|
|
except json.JSONDecodeError as exc:
|
||
|
|
raise DataError(f"{path}:{line_number}: invalid JSON: {exc}") from exc
|
||
|
|
if not isinstance(value, dict):
|
||
|
|
raise DataError(f"{path}:{line_number}: expected a JSON object")
|
||
|
|
records.append(value)
|
||
|
|
return records
|
||
|
|
|
||
|
|
|
||
|
|
def validate_source_record(record: dict[str, Any], location: str) -> None:
|
||
|
|
keys = set(record)
|
||
|
|
if keys != SOURCE_FIELDS:
|
||
|
|
missing = sorted(SOURCE_FIELDS - keys)
|
||
|
|
extra = sorted(keys - SOURCE_FIELDS)
|
||
|
|
details = []
|
||
|
|
if missing:
|
||
|
|
details.append(f"missing {', '.join(missing)}")
|
||
|
|
if extra:
|
||
|
|
details.append(f"unexpected {', '.join(extra)}")
|
||
|
|
raise DataError(f"{location}: invalid fields ({'; '.join(details)})")
|
||
|
|
|
||
|
|
prompt = record["prompt"]
|
||
|
|
if not isinstance(prompt, str) or not prompt.strip():
|
||
|
|
raise DataError(f"{location}: prompt must be a non-empty string")
|
||
|
|
if record["purpose"] not in LABEL_SET:
|
||
|
|
raise DataError(f"{location}: invalid purpose {record['purpose']!r}")
|
||
|
|
|
||
|
|
secondary = record["secondary"]
|
||
|
|
mixed = record["mixed"]
|
||
|
|
if type(mixed) is not bool:
|
||
|
|
raise DataError(f"{location}: mixed must be a boolean")
|
||
|
|
if secondary is not None and secondary not in LABEL_SET:
|
||
|
|
raise DataError(f"{location}: invalid secondary purpose {secondary!r}")
|
||
|
|
if mixed != (secondary is not None):
|
||
|
|
raise DataError(f"{location}: mixed and secondary disagree")
|
||
|
|
if secondary == record["purpose"]:
|
||
|
|
raise DataError(f"{location}: secondary must differ from purpose")
|
||
|
|
|
||
|
|
difficulty = record["difficulty"]
|
||
|
|
if (
|
||
|
|
isinstance(difficulty, bool)
|
||
|
|
or not isinstance(difficulty, (int, float))
|
||
|
|
or not math.isfinite(difficulty)
|
||
|
|
or not 0.0 <= difficulty <= 1.0
|
||
|
|
):
|
||
|
|
raise DataError(f"{location}: difficulty must be a finite value from 0 to 1")
|
||
|
|
if record["slice"] not in SLICES:
|
||
|
|
raise DataError(f"{location}: invalid slice {record['slice']!r}")
|
||
|
|
if (record["slice"] == "mixed") != mixed:
|
||
|
|
raise DataError(f"{location}: the mixed slice and mixed field disagree")
|
||
|
|
if not isinstance(record["lang"], str) or not record["lang"]:
|
||
|
|
raise DataError(f"{location}: lang must be a non-empty string")
|
||
|
|
|
||
|
|
|
||
|
|
def load_sources(paths: Sequence[Path]) -> list[SourceRecord]:
|
||
|
|
records: list[SourceRecord] = []
|
||
|
|
for path in paths:
|
||
|
|
for line, value in enumerate(load_jsonl(path), 1):
|
||
|
|
validate_source_record(value, f"{path}:{line}")
|
||
|
|
records.append(SourceRecord(value=value, source=path, line=line))
|
||
|
|
return records
|
||
|
|
|
||
|
|
|
||
|
|
def load_classifiable_fixtures(path: Path) -> list[dict[str, str]]:
|
||
|
|
try:
|
||
|
|
value = json.loads(path.read_text(encoding="utf-8"))
|
||
|
|
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||
|
|
raise DataError(f"{path}: cannot load fixtures: {exc}") from exc
|
||
|
|
if not isinstance(value, list):
|
||
|
|
raise DataError(f"{path}: fixture root must be an array")
|
||
|
|
|
||
|
|
fixtures: list[dict[str, str]] = []
|
||
|
|
for index, item in enumerate(value):
|
||
|
|
if not isinstance(item, dict):
|
||
|
|
raise DataError(f"{path}: fixture {index} must be an object")
|
||
|
|
prompt = item.get("prompt")
|
||
|
|
purpose = item.get("purpose")
|
||
|
|
if not isinstance(prompt, str) or not prompt.strip():
|
||
|
|
raise DataError(f"{path}: fixture {index} has an invalid prompt")
|
||
|
|
if purpose == "general":
|
||
|
|
continue
|
||
|
|
if purpose not in LABEL_SET:
|
||
|
|
raise DataError(f"{path}: fixture {index} has invalid purpose {purpose!r}")
|
||
|
|
fixtures.append({"prompt": prompt, "purpose": purpose})
|
||
|
|
return fixtures
|
||
|
|
|
||
|
|
|
||
|
|
def _word_shingles(prompt: str) -> frozenset[str]:
|
||
|
|
words = WORD_RE.findall(normalized_key(prompt))
|
||
|
|
if len(words) < 8:
|
||
|
|
return frozenset()
|
||
|
|
return frozenset(
|
||
|
|
"\x1f".join(words[index : index + 3])
|
||
|
|
for index in range(len(words) - 2)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _simhash(features: frozenset[str]) -> int:
|
||
|
|
weights = [0] * 64
|
||
|
|
for feature in features:
|
||
|
|
value = int.from_bytes(
|
||
|
|
hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest(), "big"
|
||
|
|
)
|
||
|
|
for bit in range(64):
|
||
|
|
weights[bit] += 1 if value & (1 << bit) else -1
|
||
|
|
|
||
|
|
signature = 0
|
||
|
|
for bit, weight in enumerate(weights):
|
||
|
|
if weight >= 0:
|
||
|
|
signature |= 1 << bit
|
||
|
|
return signature
|
||
|
|
|
||
|
|
|
||
|
|
class _NearDuplicateIndex:
|
||
|
|
"""Small dependency-free LSH pass for generated template duplicates.
|
||
|
|
|
||
|
|
This is intentionally conservative. Short prompts are handled only by exact-match
|
||
|
|
checks because a one-word change can completely change their purpose. Longer prompts
|
||
|
|
become word-trigram sets. SimHash bands produce candidates; exact Jaccard similarity
|
||
|
|
decides whether a row is dropped.
|
||
|
|
"""
|
||
|
|
|
||
|
|
def __init__(self, threshold: float) -> None:
|
||
|
|
self.threshold = threshold
|
||
|
|
self.items: list[tuple[frozenset[str], str, str]] = []
|
||
|
|
self.buckets: dict[tuple[int, int], list[int]] = defaultdict(list)
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def _bands(signature: int) -> Iterable[tuple[int, int]]:
|
||
|
|
# Requiring two matching 8-bit bands keeps random candidate sets small while
|
||
|
|
# retaining every pair whose SimHashes differ in at most six bands.
|
||
|
|
for band in range(8):
|
||
|
|
yield band, (signature >> (band * 8)) & 0xFF
|
||
|
|
|
||
|
|
def find(self, prompt: str) -> tuple[str, str, float] | None:
|
||
|
|
features = _word_shingles(prompt)
|
||
|
|
if not features:
|
||
|
|
return None
|
||
|
|
signature = _simhash(features)
|
||
|
|
hits: Counter[int] = Counter()
|
||
|
|
for band in self._bands(signature):
|
||
|
|
hits.update(self.buckets.get(band, ()))
|
||
|
|
|
||
|
|
best: tuple[str, str, float] | None = None
|
||
|
|
for index, matching_bands in hits.items():
|
||
|
|
if matching_bands < 2:
|
||
|
|
continue
|
||
|
|
other_features, other_label, other_hash = self.items[index]
|
||
|
|
union = len(features | other_features)
|
||
|
|
similarity = len(features & other_features) / union if union else 1.0
|
||
|
|
if similarity >= self.threshold and (
|
||
|
|
best is None or similarity > best[2]
|
||
|
|
):
|
||
|
|
best = (other_label, other_hash, similarity)
|
||
|
|
return best
|
||
|
|
|
||
|
|
def add(self, prompt: str, label: str) -> None:
|
||
|
|
features = _word_shingles(prompt)
|
||
|
|
if not features:
|
||
|
|
return
|
||
|
|
signature = _simhash(features)
|
||
|
|
index = len(self.items)
|
||
|
|
self.items.append((features, label, prompt_hash(prompt)))
|
||
|
|
for band in self._bands(signature):
|
||
|
|
self.buckets[band].append(index)
|
||
|
|
|
||
|
|
|
||
|
|
def curate_records(
|
||
|
|
records: Sequence[SourceRecord],
|
||
|
|
fixtures: Sequence[dict[str, str]],
|
||
|
|
*,
|
||
|
|
near_duplicate_threshold: float = 0.92,
|
||
|
|
) -> CurationResult:
|
||
|
|
if not 0.0 < near_duplicate_threshold <= 1.0:
|
||
|
|
raise DataError("near-duplicate threshold must be in (0, 1]")
|
||
|
|
|
||
|
|
exact: dict[str, tuple[str, str]] = {}
|
||
|
|
near = _NearDuplicateIndex(near_duplicate_threshold)
|
||
|
|
for fixture in fixtures:
|
||
|
|
key = normalized_key(fixture["prompt"])
|
||
|
|
previous = exact.get(key)
|
||
|
|
if previous is not None and previous[0] != fixture["purpose"]:
|
||
|
|
raise DataError("shipped fixtures contain an exact prompt with two labels")
|
||
|
|
exact[key] = (fixture["purpose"], prompt_hash(fixture["prompt"]))
|
||
|
|
near.add(fixture["prompt"], fixture["purpose"])
|
||
|
|
|
||
|
|
kept: list[SourceRecord] = []
|
||
|
|
duplicates: list[Duplicate] = []
|
||
|
|
conflicts: list[str] = []
|
||
|
|
for record in records:
|
||
|
|
prompt = record.value["prompt"]
|
||
|
|
purpose = record.value["purpose"]
|
||
|
|
key = normalized_key(prompt)
|
||
|
|
previous = exact.get(key)
|
||
|
|
if previous is not None:
|
||
|
|
previous_label, previous_hash = previous
|
||
|
|
if previous_label != purpose:
|
||
|
|
conflicts.append(
|
||
|
|
f"{record.source}:{record.line}: exact duplicate has labels "
|
||
|
|
f"{previous_label!r} and {purpose!r}"
|
||
|
|
)
|
||
|
|
else:
|
||
|
|
duplicates.append(
|
||
|
|
Duplicate(record, previous_hash, "exact", 1.0)
|
||
|
|
)
|
||
|
|
continue
|
||
|
|
|
||
|
|
match = near.find(prompt)
|
||
|
|
if match is not None:
|
||
|
|
previous_label, previous_hash, similarity = match
|
||
|
|
if previous_label != purpose:
|
||
|
|
conflicts.append(
|
||
|
|
f"{record.source}:{record.line}: {similarity:.1%}-similar prompt "
|
||
|
|
f"has labels {previous_label!r} and {purpose!r}"
|
||
|
|
)
|
||
|
|
# Keep the row for now so all conflicts are reported without causing a
|
||
|
|
# cascade of duplicates against a record that may later be relabeled.
|
||
|
|
exact[key] = (purpose, prompt_hash(prompt))
|
||
|
|
near.add(prompt, purpose)
|
||
|
|
kept.append(record)
|
||
|
|
else:
|
||
|
|
duplicates.append(
|
||
|
|
Duplicate(record, previous_hash, "near", similarity)
|
||
|
|
)
|
||
|
|
continue
|
||
|
|
|
||
|
|
exact[key] = (purpose, prompt_hash(prompt))
|
||
|
|
near.add(prompt, purpose)
|
||
|
|
kept.append(record)
|
||
|
|
|
||
|
|
if conflicts:
|
||
|
|
preview = "\n".join(conflicts[:20])
|
||
|
|
remainder = len(conflicts) - min(20, len(conflicts))
|
||
|
|
suffix = f"\n... {remainder} more conflict(s)" if remainder else ""
|
||
|
|
raise DataError(f"near-duplicate label conflicts require review:\n{preview}{suffix}")
|
||
|
|
return CurationResult(records=kept, duplicates=duplicates)
|
||
|
|
|
||
|
|
|
||
|
|
def _stable_digest(seed: int, prompt: str) -> str:
|
||
|
|
material = f"{seed}\0{normalized_key(prompt)}".encode("utf-8")
|
||
|
|
return hashlib.sha256(material).hexdigest()
|
||
|
|
|
||
|
|
|
||
|
|
def _stratified_order(
|
||
|
|
records: Sequence[SourceRecord],
|
||
|
|
*,
|
||
|
|
seed: int,
|
||
|
|
strata: Callable[[SourceRecord], tuple[str, ...]],
|
||
|
|
) -> list[SourceRecord]:
|
||
|
|
groups: dict[tuple[str, ...], list[SourceRecord]] = defaultdict(list)
|
||
|
|
for record in records:
|
||
|
|
groups[strata(record)].append(record)
|
||
|
|
|
||
|
|
ranked: list[tuple[float, str, SourceRecord]] = []
|
||
|
|
for key in sorted(groups):
|
||
|
|
group = sorted(
|
||
|
|
groups[key],
|
||
|
|
key=lambda record: _stable_digest(seed, record.value["prompt"]),
|
||
|
|
)
|
||
|
|
size = len(group)
|
||
|
|
for index, record in enumerate(group):
|
||
|
|
quantile = (index + 0.5) / size
|
||
|
|
ranked.append(
|
||
|
|
(quantile, _stable_digest(seed + 1, record.value["prompt"]), record)
|
||
|
|
)
|
||
|
|
return [item[2] for item in sorted(ranked, key=lambda item: (item[0], item[1]))]
|
||
|
|
|
||
|
|
|
||
|
|
def split_records(
|
||
|
|
records: Sequence[SourceRecord],
|
||
|
|
*,
|
||
|
|
fixture_count: int,
|
||
|
|
seed: int = 0xC1A551F1,
|
||
|
|
train_ratio: float = 0.8,
|
||
|
|
validation_ratio: float = 0.1,
|
||
|
|
) -> SplitResult:
|
||
|
|
if not records:
|
||
|
|
raise DataError("cannot split an empty dataset")
|
||
|
|
if fixture_count < 0:
|
||
|
|
raise DataError("fixture_count cannot be negative")
|
||
|
|
if not 0.0 < train_ratio < 1.0 or not 0.0 < validation_ratio < 1.0:
|
||
|
|
raise DataError("split ratios must be in (0, 1)")
|
||
|
|
if train_ratio + validation_ratio >= 1.0:
|
||
|
|
raise DataError("train + validation ratios must leave room for test")
|
||
|
|
|
||
|
|
logical_total = len(records) + fixture_count
|
||
|
|
target_train = round(logical_total * train_ratio)
|
||
|
|
target_validation = round(logical_total * validation_ratio)
|
||
|
|
target_test = logical_total - target_train - target_validation
|
||
|
|
if fixture_count > target_test:
|
||
|
|
raise DataError("fixture count exceeds the target test split")
|
||
|
|
|
||
|
|
vague = [record for record in records if record.value["slice"] == "vague-eval"]
|
||
|
|
regular = [record for record in records if record.value["slice"] != "vague-eval"]
|
||
|
|
eval_capacity = target_validation + target_test - fixture_count
|
||
|
|
if len(vague) > eval_capacity:
|
||
|
|
raise DataError(
|
||
|
|
"vague-eval records exceed validation/test capacity; lower train_ratio"
|
||
|
|
)
|
||
|
|
|
||
|
|
# Allocate vague records between validation and test in proportion to each split's
|
||
|
|
# remaining capacity. None may enter training.
|
||
|
|
test_source_capacity = target_test - fixture_count
|
||
|
|
vague_validation_count = round(
|
||
|
|
len(vague) * target_validation / (target_validation + test_source_capacity)
|
||
|
|
)
|
||
|
|
vague_validation_count = min(vague_validation_count, target_validation)
|
||
|
|
vague_test_count = len(vague) - vague_validation_count
|
||
|
|
if vague_test_count > test_source_capacity:
|
||
|
|
overflow = vague_test_count - test_source_capacity
|
||
|
|
vague_validation_count += overflow
|
||
|
|
vague_test_count -= overflow
|
||
|
|
|
||
|
|
vague_order = _stratified_order(
|
||
|
|
vague,
|
||
|
|
seed=seed + 7,
|
||
|
|
strata=lambda record: (record.value["purpose"],),
|
||
|
|
)
|
||
|
|
vague_validation = vague_order[:vague_validation_count]
|
||
|
|
vague_test = vague_order[vague_validation_count:]
|
||
|
|
|
||
|
|
regular_train_count = target_train
|
||
|
|
regular_validation_count = target_validation - len(vague_validation)
|
||
|
|
regular_test_count = test_source_capacity - len(vague_test)
|
||
|
|
if (
|
||
|
|
regular_train_count + regular_validation_count + regular_test_count
|
||
|
|
!= len(regular)
|
||
|
|
):
|
||
|
|
raise DataError("internal split accounting mismatch")
|
||
|
|
|
||
|
|
regular_order = _stratified_order(
|
||
|
|
regular,
|
||
|
|
seed=seed,
|
||
|
|
strata=lambda record: (
|
||
|
|
record.value["purpose"],
|
||
|
|
record.value["slice"],
|
||
|
|
record.value["lang"].split("-", 1)[0].casefold()
|
||
|
|
if record.value["lang"].split("-", 1)[0].casefold() != "en"
|
||
|
|
else "en",
|
||
|
|
),
|
||
|
|
)
|
||
|
|
train = regular_order[:regular_train_count]
|
||
|
|
validation_end = regular_train_count + regular_validation_count
|
||
|
|
validation = regular_order[regular_train_count:validation_end] + vague_validation
|
||
|
|
test = regular_order[validation_end:] + vague_test
|
||
|
|
|
||
|
|
# A second stable ordering makes file content independent of stratum dictionary order
|
||
|
|
# and gives the trainer a deterministic shuffle before its epoch sampler takes over.
|
||
|
|
def stable_sort(items: Sequence[SourceRecord], offset: int) -> list[SourceRecord]:
|
||
|
|
return sorted(
|
||
|
|
items,
|
||
|
|
key=lambda record: _stable_digest(seed + offset, record.value["prompt"]),
|
||
|
|
)
|
||
|
|
|
||
|
|
result = SplitResult(
|
||
|
|
train=stable_sort(train, 11),
|
||
|
|
validation=stable_sort(validation, 13),
|
||
|
|
test=stable_sort(test, 17),
|
||
|
|
fixture_count=fixture_count,
|
||
|
|
)
|
||
|
|
if any(record.value["slice"] == "vague-eval" for record in result.train):
|
||
|
|
raise DataError("vague-eval leakage into training")
|
||
|
|
if len(result.train) != target_train:
|
||
|
|
raise DataError("train split missed its target size")
|
||
|
|
if len(result.validation) != target_validation:
|
||
|
|
raise DataError("validation split missed its target size")
|
||
|
|
if result.logical_test_count != target_test:
|
||
|
|
raise DataError("test split missed its target size")
|
||
|
|
return result
|
||
|
|
|
||
|
|
|
||
|
|
def distribution(records: Sequence[SourceRecord]) -> dict[str, dict[str, int]]:
|
||
|
|
return {
|
||
|
|
"purpose": dict(
|
||
|
|
sorted(Counter(record.value["purpose"] for record in records).items())
|
||
|
|
),
|
||
|
|
"slice": dict(
|
||
|
|
sorted(Counter(record.value["slice"] for record in records).items())
|
||
|
|
),
|
||
|
|
"language": dict(
|
||
|
|
sorted(
|
||
|
|
Counter(
|
||
|
|
record.value["lang"].split("-", 1)[0].casefold()
|
||
|
|
for record in records
|
||
|
|
).items()
|
||
|
|
)
|
||
|
|
),
|
||
|
|
}
|