Files
nucleic-purpose-classifier/purpose_data.py
T

584 lines
21 KiB
Python
Raw Normal View History

"""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 exclude_reviewed_duplicates(
records: Sequence[SourceRecord],
decisions: Sequence[dict[str, Any]],
) -> CurationResult:
"""Apply human-reviewed semantic exclusions, failing closed on corpus drift."""
by_hash = {prompt_hash(record.value["prompt"]): record for record in records}
if len(by_hash) != len(records):
raise DataError("reviewed exclusions require an exact-deduplicated population")
drops: dict[str, tuple[str, float]] = {}
for index, decision in enumerate(decisions, 1):
try:
dropped_hash = decision["droppedPromptHash"]
matched_hash = decision["matchedPromptHash"]
similarity = float(decision["similarity"])
except (KeyError, TypeError, ValueError) as exc:
raise DataError(
f"semantic exclusion {index}: invalid decision fields"
) from exc
if (
not isinstance(dropped_hash, str)
or not isinstance(matched_hash, str)
or len(dropped_hash) != 64
or len(matched_hash) != 64
or not 0.0 <= similarity <= 1.0
):
raise DataError(f"semantic exclusion {index}: invalid hashes/similarity")
if dropped_hash == matched_hash:
raise DataError(f"semantic exclusion {index}: cannot match itself")
if dropped_hash in drops:
raise DataError(f"semantic exclusion {index}: duplicate dropped hash")
drops[dropped_hash] = (matched_hash, similarity)
missing = sorted((set(drops) | {item[0] for item in drops.values()}) - set(by_hash))
if missing:
raise DataError(
"reviewed semantic exclusion no longer matches the curated corpus: "
+ ", ".join(missing)
)
kept = []
duplicates = []
for record in records:
dropped_hash = prompt_hash(record.value["prompt"])
decision = drops.get(dropped_hash)
if decision is None:
kept.append(record)
continue
matched_hash, similarity = decision
matched = by_hash[matched_hash]
if record.value["purpose"] != matched.value["purpose"]:
raise DataError(
"reviewed semantic duplicate labels no longer agree: "
f"{dropped_hash} vs {matched_hash}"
)
duplicates.append(
Duplicate(
dropped=record,
matched_prompt_hash=matched_hash,
kind="semantic-reviewed",
similarity=similarity,
)
)
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()
)
),
}