Files
nucleic-purpose-classifier/validate-data.py
T

767 lines
26 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""Validate purpose-classifier JSONL data.
The checks in this script are derived from datagen-prompt.md. Definite format and
process violations are errors. Approximate targets (the requirements written as
"~" or "≈") are warnings; pass --strict to make warnings fail the command too.
Usage:
python3 ml/purpose-classifier/validate-data.py
python3 ml/purpose-classifier/validate-data.py path/to/batch.jsonl --batch-size 200
python3 ml/purpose-classifier/validate-data.py --strict
"""
from __future__ import annotations
import argparse
import json
import math
import re
import sys
import unicodedata
from collections import Counter
from dataclasses import dataclass
from datetime import date
from pathlib import Path
from typing import Any, Iterable
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATA_DIR = SCRIPT_DIR / "data"
DEFAULT_CANONICAL_FILES = (
DEFAULT_DATA_DIR / "purpose-prompts.jsonl",
DEFAULT_DATA_DIR / "purpose-prompts-round2.jsonl",
)
DEFAULT_FIXTURES = (
SCRIPT_DIR.parent.parent
/ "Tests"
/ "NucleicCoreTests"
/ "Fixtures"
/ "purpose-prompts.json"
)
FIELDS = frozenset(
{"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"}
)
PURPOSES = (
"planning",
"backendImpl",
"frontendImpl",
"quickFix",
"refactor",
"debugging",
"review",
"writing",
)
PURPOSE_SET = frozenset(PURPOSES)
SLICES = ("core", "boundary", "mixed", "pasted-context", "vague-eval")
SLICE_SET = frozenset(SLICES)
SLICE_TARGETS = {
"core": 0.55,
"boundary": 0.20,
"mixed": 0.10,
"pasted-context": 0.10,
"vague-eval": 0.05,
}
EXPECTED_NON_ENGLISH = frozenset({"es", "de", "fr", "pt", "zh", "ja"})
# This is deliberately a conservative list. It is only used for the prompt's
# mechanical "same opening verb no more than three times per 50 examples" check;
# unknown first words are not guessed to be verbs.
OPENING_VERBS = frozenset(
{
"add",
"analyze",
"architect",
"audit",
"break",
"build",
"bump",
"change",
"check",
"clean",
"compare",
"consolidate",
"correct",
"create",
"debug",
"decouple",
"design",
"diagnose",
"document",
"draft",
"evaluate",
"explain",
"extract",
"figure",
"find",
"fix",
"flip",
"help",
"implement",
"investigate",
"look",
"make",
"map",
"migrate",
"modularize",
"move",
"outline",
"plan",
"polish",
"proofread",
"propose",
"refactor",
"remove",
"rename",
"replace",
"restructure",
"review",
"rewrite",
"set",
"simplify",
"sketch",
"split",
"summarize",
"track",
"translate",
"tweak",
"update",
"walk",
"wire",
"write",
}
)
TOKEN_RE = re.compile(r"\w+|[^\w\s]", re.UNICODE)
WORD_RE = re.compile(r"[^\W_]+(?:['’-][^\W_]+)*", re.UNICODE)
ENUMERATED_TASK_RE = re.compile(r"^\s*(?:[-*]\s*)?task\s+\d+\s*:", re.IGNORECASE)
LABEL_LEAK_RE = re.compile(
r"\b(?:this|it)\s+is\s+(?:an?\s+)?"
r"(?:planning|backendimpl|frontendimpl|quickfix|refactor|debugging|review|writing)"
r"\s+(?:task|prompt)\b",
re.IGNORECASE,
)
META_RE = re.compile(r"\b(?:classify this prompt|as an ai)\b", re.IGNORECASE)
MID_CONVERSATION_RE = re.compile(
r"^\s*(?:"
r"yes[,\s]+do\s+option\s+\d+|"
r"that\s+didn['’]?t\s+work|"
r"try\s+again\b|"
r"same\s+error\s+as\s+before|"
r"looks\s+good[,\s]+ship\s+it|"
r"no[,\s]+the\s+other\b"
r")",
re.IGNORECASE,
)
@dataclass(frozen=True)
class Issue:
severity: str
path: Path
message: str
line: int | None = None
def render(self) -> str:
location = str(self.path)
if self.line is not None:
location += f":{self.line}"
return f"{location}: {self.severity}: {self.message}"
@dataclass(frozen=True)
class Record:
path: Path
line: int
value: dict[str, Any]
class Validator:
def __init__(
self,
*,
batch_size: int = 0,
expected_total: int = 12_214,
fixture_path: Path | None = DEFAULT_FIXTURES,
process_checks: bool = True,
) -> None:
self.batch_size = batch_size
self.expected_total = expected_total
self.fixture_path = fixture_path
self.process_checks = process_checks
self.issues: list[Issue] = []
def error(self, path: Path, message: str, line: int | None = None) -> None:
self.issues.append(Issue("error", path, message, line))
def warning(self, path: Path, message: str, line: int | None = None) -> None:
self.issues.append(Issue("warning", path, message, line))
def validate(self, paths: list[Path], roots: list[Path]) -> list[Record]:
records_by_file: dict[Path, list[Record]] = {}
for path in paths:
records_by_file[path] = self._read_jsonl(path)
records = [record for path in paths for record in records_by_file[path]]
self._check_duplicates(records)
self._check_fixture_contamination(records)
if self.process_checks:
for path, file_records in records_by_file.items():
self._check_file_batches(path, file_records)
self._check_global_distribution(records, roots)
self._check_manifests(roots)
return records
def _read_jsonl(self, path: Path) -> list[Record]:
records: list[Record] = []
try:
lines = path.read_text(encoding="utf-8").splitlines()
except (OSError, UnicodeError) as exc:
self.error(path, f"cannot read UTF-8 JSONL: {exc}")
return records
if not lines:
self.error(path, "file is empty")
return records
for line_number, line in enumerate(lines, 1):
if not line.strip():
self.error(path, "blank lines are not allowed in strict JSONL", line_number)
continue
try:
value = json.loads(line, parse_constant=self._reject_json_constant)
except (json.JSONDecodeError, ValueError) as exc:
self.error(path, f"invalid JSON: {exc}", line_number)
continue
if not isinstance(value, dict):
self.error(path, "each JSONL line must be an object", line_number)
continue
if self._check_record(path, line_number, value):
records.append(Record(path, line_number, value))
return records
@staticmethod
def _reject_json_constant(value: str) -> None:
raise ValueError(f"{value} is not valid strict JSON")
def _check_record(self, path: Path, line: int, value: dict[str, Any]) -> bool:
valid = True
keys = set(value)
missing = sorted(FIELDS - keys)
extra = sorted(keys - FIELDS)
if missing:
self.error(path, f"missing fields: {', '.join(missing)}", line)
valid = False
if extra:
self.error(path, f"unexpected fields: {', '.join(extra)}", line)
valid = False
if missing:
return False
prompt = value["prompt"]
if not isinstance(prompt, str):
self.error(path, "prompt must be a string", line)
valid = False
elif not prompt.strip():
self.error(path, "prompt must not be empty or whitespace-only", line)
valid = False
purpose = value["purpose"]
if not isinstance(purpose, str) or purpose not in PURPOSE_SET:
self.error(path, f"purpose must be one of: {', '.join(PURPOSES)}", line)
valid = False
secondary = value["secondary"]
if secondary is not None and (
not isinstance(secondary, str) or secondary not in PURPOSE_SET
):
self.error(path, "secondary must be null or a valid purpose label", line)
valid = False
mixed = value["mixed"]
if type(mixed) is not bool:
self.error(path, "mixed must be a boolean", line)
valid = False
difficulty = value["difficulty"]
if (
isinstance(difficulty, bool)
or not isinstance(difficulty, (int, float))
or not math.isfinite(difficulty)
or not 0.0 <= difficulty <= 1.0
):
self.error(path, "difficulty must be a finite number from 0.0 to 1.0", line)
valid = False
slice_name = value["slice"]
if not isinstance(slice_name, str) or slice_name not in SLICE_SET:
self.error(path, f"slice must be one of: {', '.join(SLICES)}", line)
valid = False
lang = value["lang"]
if not isinstance(lang, str) or not is_bcp47(lang):
self.error(path, "lang must be a syntactically valid BCP 47 tag", line)
valid = False
if type(mixed) is bool:
if mixed and secondary is None:
self.error(path, "mixed=true requires a secondary purpose", line)
valid = False
if not mixed and secondary is not None:
self.error(path, "mixed=false requires secondary=null", line)
valid = False
if (slice_name == "mixed") != mixed:
self.error(path, "the mixed slice and mixed field must agree", line)
valid = False
if purpose in PURPOSE_SET and secondary == purpose:
self.error(path, "secondary must differ from the primary purpose", line)
valid = False
if isinstance(prompt, str) and prompt.strip():
self._check_prompt_antipatterns(path, line, prompt)
tokens = token_count(prompt)
if slice_name == "pasted-context" and not 100 <= tokens <= 400:
self.warning(
path,
"pasted-context prompt should be approximately 100-400 tokens "
f"(estimated {tokens})",
line,
)
elif tokens > 400:
self.warning(
path,
f"prompt exceeds the approximately 400-token maximum (estimated {tokens})",
line,
)
return valid
def _check_prompt_antipatterns(self, path: Path, line: int, prompt: str) -> None:
if ENUMERATED_TASK_RE.search(prompt):
self.error(path, "enumerated 'Task N:' prefix is forbidden", line)
if LABEL_LEAK_RE.search(prompt):
self.error(path, "prompt leaks its dataset label as a task/prompt hint", line)
if META_RE.search(prompt):
self.error(path, "assistant-directed classification/meta prompt is forbidden", line)
if MID_CONVERSATION_RE.search(prompt):
self.error(path, "prompt reads as a mid-conversation reply", line)
def _check_duplicates(self, records: list[Record]) -> None:
seen: dict[str, Record] = {}
for record in records:
key = normalize_prompt(record.value["prompt"])
previous = seen.get(key)
if previous is not None:
self.error(
record.path,
f"duplicate prompt; first seen at {previous.path}:{previous.line}",
record.line,
)
else:
seen[key] = record
def _check_fixture_contamination(self, records: list[Record]) -> None:
if self.fixture_path is None:
return
try:
fixture_data = json.loads(self.fixture_path.read_text(encoding="utf-8"))
fixture_prompts = {
normalize_prompt(item["prompt"])
for item in fixture_data
if isinstance(item, dict) and isinstance(item.get("prompt"), str)
}
except (OSError, UnicodeError, json.JSONDecodeError, TypeError, KeyError) as exc:
self.warning(self.fixture_path, f"could not load eval-only fixtures: {exc}")
return
for record in records:
if normalize_prompt(record.value["prompt"]) in fixture_prompts:
self.error(
record.path,
"prompt duplicates an eval-only purpose-prompts.json fixture",
record.line,
)
def _check_file_batches(self, path: Path, records: list[Record]) -> None:
if self.batch_size <= 0:
return
if len(records) % self.batch_size:
self.error(
path,
f"contains {len(records)} valid records; generation batches must contain "
f"{self.batch_size} records",
)
for start in range(0, len(records), self.batch_size):
batch = records[start : start + self.batch_size]
if batch:
self._check_batch(path, start // self.batch_size + 1, batch)
def _check_batch(self, path: Path, number: int, records: list[Record]) -> None:
name = f"batch {number}"
total = len(records)
minimum_length_count = math.ceil(total * 0.15)
short_count = sum(word_count(r.value["prompt"]) < 8 for r in records)
long_count = sum(token_count(r.value["prompt"]) > 60 for r in records)
if short_count < minimum_length_count:
self.error(
path,
f"{name} has {short_count}/{total} prompts under 8 words; "
f"at least {minimum_length_count} required",
)
if long_count < minimum_length_count:
self.error(
path,
f"{name} has {long_count}/{total} prompts over 60 estimated tokens; "
f"at least {minimum_length_count} required",
)
slices = Counter(r.value["slice"] for r in records)
for slice_name, target in SLICE_TARGETS.items():
actual = slices[slice_name] / total
if abs(actual - target) > 0.05:
self.warning(
path,
f"{name} {slice_name} share is {actual:.1%}; target is approximately "
f"{target:.0%} (±5 percentage points)",
)
mixed_share = sum(r.value["mixed"] for r in records) / total
if mixed_share > 0.12:
self.error(
path,
f"{name} mixed-intent share is {mixed_share:.1%}; maximum is approximately 12%",
)
english = sum(primary_language(r.value["lang"]) == "en" for r in records)
english_share = english / total
if not 0.90 <= english_share <= 0.98:
self.warning(
path,
f"{name} English share is {english_share:.1%}; target is approximately 95%",
)
foreign_languages = {
primary_language(r.value["lang"])
for r in records
if primary_language(r.value["lang"]) != "en"
}
unexpected = sorted(foreign_languages - EXPECTED_NON_ENGLISH)
if unexpected:
self.warning(
path,
f"{name} uses non-English languages outside the requested set: "
f"{', '.join(unexpected)}",
)
for window_start in range(0, total, 50):
window = records[window_start : window_start + 50]
verbs: dict[str, list[Record]] = {}
for record in window:
verb = opening_verb(record.value["prompt"])
if verb is not None:
verbs.setdefault(verb, []).append(record)
for verb, matches in sorted(verbs.items()):
if len(matches) > 3:
locations = ", ".join(str(r.line) for r in matches)
self.error(
path,
f"{name}, records {window_start + 1}-{window_start + len(window)} "
f"open with '{verb}' {len(matches)} times (lines {locations}); maximum is 3",
)
def _check_global_distribution(self, records: list[Record], roots: list[Path]) -> None:
label = roots[0] if roots else DEFAULT_DATA_DIR
total = len(records)
if self.expected_total > 0 and total != self.expected_total:
self.warning(
label,
f"dataset contains {total} valid records; generation target is "
f"{self.expected_total}",
)
if not records:
return
non_vague = [r for r in records if r.value["slice"] != "vague-eval"]
if len(non_vague) >= len(PURPOSES):
counts = Counter(r.value["purpose"] for r in non_vague)
expected = len(non_vague) / len(PURPOSES)
low = expected * 0.85
high = expected * 1.15
outside = [
f"{purpose}={counts[purpose]}"
for purpose in PURPOSES
if not low <= counts[purpose] <= high
]
if outside:
self.error(
label,
"non-vague primary labels are not within ±15% of uniform "
f"(expected about {expected:.1f} each): {', '.join(outside)}",
)
mixed_share = sum(r.value["mixed"] for r in records) / total
if mixed_share > 0.12:
self.error(
label,
f"dataset mixed-intent share is {mixed_share:.1%}; maximum is approximately 12%",
)
slices = Counter(r.value["slice"] for r in records)
for slice_name, target in SLICE_TARGETS.items():
actual = slices[slice_name] / total
if abs(actual - target) > 0.05:
self.warning(
label,
f"dataset {slice_name} share is {actual:.1%}; target is approximately "
f"{target:.0%} (±5 percentage points)",
)
def _check_manifests(self, roots: list[Path]) -> None:
directories = sorted({root if root.is_dir() else root.parent for root in roots})
for directory in directories:
candidates = sorted(directory.glob("*generation-manifest*.json"))
if not candidates:
self.warning(
directory,
"generation manifest not found; record the model, date, and batch topics",
)
continue
for candidate in candidates:
self._check_manifest(candidate)
def _check_manifest(self, path: Path) -> None:
try:
value = json.loads(path.read_text(encoding="utf-8"))
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
self.error(path, f"invalid generation manifest: {exc}")
return
if not isinstance(value, dict):
self.error(path, "generation manifest must be a JSON object")
return
values_by_key: dict[str, list[Any]] = {}
for key, item in walk_mapping_items(value):
normalized_key = re.sub(r"[-_]", "", key).casefold()
values_by_key.setdefault(normalized_key, []).append(item)
models = values_by_key.get("model", []) + values_by_key.get("generatingmodel", [])
if not any(isinstance(item, str) and item.strip() for item in models):
self.error(path, "generation manifest must record a non-empty model")
dates = values_by_key.get("date", []) + values_by_key.get("generationdate", [])
valid_dates = [item for item in dates if isinstance(item, str) and is_iso_date(item)]
if not valid_dates:
self.error(path, "generation manifest must record an ISO date (YYYY-MM-DD)")
topics = values_by_key.get("topics", []) + values_by_key.get("batchtopics", [])
if not any(
(isinstance(item, str) and item.strip())
or (isinstance(item, list) and len(item) > 0)
for item in topics
):
self.error(path, "generation manifest must record non-empty batch topics")
def is_bcp47(value: str) -> bool:
"""A dependency-free structural BCP 47 check.
Full registry validation would make this script network- or package-dependent.
This rejects the common malformed forms while accepting normal language,
script, region, variant, extension, and private-use tags.
"""
if not value or "_" in value or value.startswith("-") or value.endswith("-"):
return False
parts = value.split("-")
if any(not part.isascii() or not part.isalnum() or not 1 <= len(part) <= 8 for part in parts):
return False
if parts[0].casefold() == "x":
return len(parts) > 1
return parts[0].isalpha() and 2 <= len(parts[0]) <= 8
def normalize_prompt(prompt: str) -> str:
normalized = unicodedata.normalize("NFKC", prompt).casefold()
return " ".join(normalized.split())
def word_count(prompt: str) -> int:
return len(WORD_RE.findall(prompt))
def token_count(prompt: str) -> int:
return len(TOKEN_RE.findall(prompt))
def primary_language(tag: str) -> str:
return tag.split("-", 1)[0].casefold()
def opening_verb(prompt: str) -> str | None:
words = [word.casefold() for word in WORD_RE.findall(prompt[:160])]
if not words:
return None
if words[0] in OPENING_VERBS:
return words[0]
# Recognize common request wrappers without mistaking a later noun ("a small
# change") for the sentence's opening verb.
start = 0
if words[0] in {"please", "kindly"}:
start = 1
elif len(words) >= 2 and words[0] in {"can", "could", "would", "will"}:
start = 2 if words[1] == "you" else 1
elif len(words) >= 2 and words[0] in {"i", "we"} and words[1] in {
"need",
"want",
"would",
}:
start = 2
elif len(words) >= 2 and words[0] == "help" and words[1] in {"me", "us"}:
start = 2
else:
return None
for word in words[start : start + 3]:
if word in OPENING_VERBS:
return word
return None
def walk_mapping_items(value: Any) -> Iterable[tuple[str, Any]]:
if isinstance(value, dict):
for key, item in value.items():
if isinstance(key, str):
yield key, item
yield from walk_mapping_items(item)
elif isinstance(value, list):
for item in value:
yield from walk_mapping_items(item)
def is_iso_date(value: str) -> bool:
try:
date.fromisoformat(value)
except ValueError:
return False
return bool(re.fullmatch(r"\d{4}-\d{2}-\d{2}", value))
def discover_paths(targets: list[Path]) -> tuple[list[Path], list[Path], list[str]]:
files: set[Path] = set()
roots: list[Path] = []
errors: list[str] = []
for target in targets:
target = target.resolve()
if not target.exists():
errors.append(f"{target}: path does not exist")
continue
roots.append(target)
if target.is_dir():
files.update(path.resolve() for path in target.rglob("*.jsonl") if path.is_file())
elif target.is_file() and target.suffix.casefold() == ".jsonl":
files.add(target)
else:
errors.append(f"{target}: expected a .jsonl file or directory")
if not files and not errors:
errors.append("no .jsonl files found")
return sorted(files), roots, errors
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Validate purpose-classifier JSONL data against datagen-prompt.md."
)
parser.add_argument(
"paths",
nargs="*",
type=Path,
help=f"JSONL files or directories (default: {DEFAULT_DATA_DIR})",
)
parser.add_argument(
"--strict",
action="store_true",
help="return failure for approximate-target warnings as well as errors",
)
parser.add_argument(
"--batch-size",
type=int,
default=0,
help="required generation batch size; use 200 for raw batches (default: disabled)",
)
parser.add_argument(
"--expected-total",
type=int,
default=12_214,
help="expected canonical record count; use 0 to disable (default: 12214)",
)
parser.add_argument(
"--fixtures",
type=Path,
default=DEFAULT_FIXTURES,
help="eval-only fixture JSON checked for contamination",
)
parser.add_argument(
"--no-fixture-check",
action="store_true",
help="do not check prompts against the eval-only fixture set",
)
parser.add_argument(
"--no-process-checks",
action="store_true",
help="only check JSONL records, schema, duplicates, and fixture contamination",
)
parser.add_argument(
"--max-issues",
type=int,
default=200,
help="maximum diagnostics to print; 0 prints all (default: 200)",
)
return parser
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
if args.batch_size < 0 or args.expected_total < 0 or args.max_issues < 0:
print("error: numeric options must be non-negative", file=sys.stderr)
return 2
targets = args.paths or list(DEFAULT_CANONICAL_FILES)
paths, roots, discovery_errors = discover_paths(targets)
if discovery_errors:
for message in discovery_errors:
print(f"error: {message}", file=sys.stderr)
return 2
validator = Validator(
batch_size=args.batch_size,
expected_total=args.expected_total,
fixture_path=None if args.no_fixture_check else args.fixtures.resolve(),
process_checks=not args.no_process_checks,
)
records = validator.validate(paths, roots)
errors = sum(issue.severity == "error" for issue in validator.issues)
warnings = sum(issue.severity == "warning" for issue in validator.issues)
visible = validator.issues if args.max_issues == 0 else validator.issues[: args.max_issues]
for issue in visible:
print(issue.render(), file=sys.stderr)
hidden = len(validator.issues) - len(visible)
if hidden:
print(f"... {hidden} additional issues omitted", file=sys.stderr)
print(
f"Validated {len(records)} records in {len(paths)} JSONL files: "
f"{errors} error(s), {warnings} warning(s)."
)
return 1 if errors or (args.strict and warnings) else 0
if __name__ == "__main__":
raise SystemExit(main())