Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,323 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Build a training-only dataset augmented with labeled first prompts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Any, Sequence
|
||||
|
||||
from purpose_data import (
|
||||
DataError,
|
||||
SourceRecord,
|
||||
curate_records,
|
||||
distribution,
|
||||
file_sha256,
|
||||
jsonl_bytes,
|
||||
load_classifiable_fixtures,
|
||||
load_jsonl,
|
||||
normalized_key,
|
||||
prompt_hash,
|
||||
validate_source_record,
|
||||
write_json,
|
||||
write_jsonl,
|
||||
)
|
||||
|
||||
|
||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||
REPOSITORY_ROOT = SCRIPT_DIR.parent.parent
|
||||
DEFAULT_BASE_DATASET = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
||||
DEFAULT_HISTORY = (
|
||||
SCRIPT_DIR / ".artifacts" / "nucleic-history-first-prompts.labeled.jsonl"
|
||||
)
|
||||
DEFAULT_FIXTURES = (
|
||||
REPOSITORY_ROOT
|
||||
/ "Tests"
|
||||
/ "NucleicCoreTests"
|
||||
/ "Fixtures"
|
||||
/ "purpose-prompts.json"
|
||||
)
|
||||
DEFAULT_OUTPUT = SCRIPT_DIR / ".artifacts" / "dataset-v1-history-first-prompts"
|
||||
EXPERIMENT_VERSION = "history-first-prompts-training-augmentation-v1"
|
||||
|
||||
|
||||
def _relative(path: Path) -> str:
|
||||
try:
|
||||
return str(path.resolve().relative_to(REPOSITORY_ROOT))
|
||||
except ValueError:
|
||||
return str(path.resolve())
|
||||
|
||||
|
||||
def _validate_records(records: Sequence[dict[str, Any]], path: Path) -> None:
|
||||
if not records:
|
||||
raise DataError(f"{path}: split is empty")
|
||||
for line, record in enumerate(records, 1):
|
||||
validate_source_record(record, f"{path}:{line}")
|
||||
|
||||
|
||||
def _indexed_records(
|
||||
records: Sequence[dict[str, Any]], path: Path
|
||||
) -> list[SourceRecord]:
|
||||
return [
|
||||
SourceRecord(value=record, source=path, line=line)
|
||||
for line, record in enumerate(records, 1)
|
||||
]
|
||||
|
||||
|
||||
def _unique_split_keys(
|
||||
splits: Sequence[tuple[str, Sequence[dict[str, Any]]]]
|
||||
) -> dict[str, tuple[str, str]]:
|
||||
seen: dict[str, tuple[str, str]] = {}
|
||||
for split_name, records in splits:
|
||||
for line, record in enumerate(records, 1):
|
||||
key = normalized_key(record["prompt"])
|
||||
previous = seen.get(key)
|
||||
if previous is not None:
|
||||
previous_split, previous_label = previous
|
||||
detail = (
|
||||
"conflicting labels"
|
||||
if previous_label != record["purpose"]
|
||||
else "duplicate prompt"
|
||||
)
|
||||
raise DataError(
|
||||
f"{split_name}:{line}: {detail} also present in {previous_split}"
|
||||
)
|
||||
seen[key] = (split_name, record["purpose"])
|
||||
return seen
|
||||
|
||||
|
||||
def _sha256_bytes(value: bytes) -> str:
|
||||
return hashlib.sha256(value).hexdigest()
|
||||
|
||||
|
||||
def _exclusion(
|
||||
record: SourceRecord,
|
||||
*,
|
||||
reason: str,
|
||||
matched_prompt_hash: str | None = None,
|
||||
similarity: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
value: dict[str, Any] = {
|
||||
"promptHash": prompt_hash(record.value["prompt"]),
|
||||
"sourceLine": record.line,
|
||||
"reason": reason,
|
||||
}
|
||||
if matched_prompt_hash is not None:
|
||||
value["matchedPromptHash"] = matched_prompt_hash
|
||||
if similarity is not None:
|
||||
value["similarity"] = round(similarity, 6)
|
||||
return value
|
||||
|
||||
|
||||
def prepare(
|
||||
*,
|
||||
base_dataset: Path,
|
||||
history_path: Path,
|
||||
fixtures_path: Path,
|
||||
output_dir: Path,
|
||||
near_duplicate_threshold: float,
|
||||
overwrite_output: bool,
|
||||
) -> dict[str, Any]:
|
||||
base_paths = {
|
||||
split: base_dataset / f"{split}.jsonl"
|
||||
for split in ("train", "validation", "test")
|
||||
}
|
||||
base = {split: load_jsonl(path) for split, path in base_paths.items()}
|
||||
for split, path in base_paths.items():
|
||||
_validate_records(base[split], path)
|
||||
|
||||
base_keys = _unique_split_keys(
|
||||
[(split, base[split]) for split in ("train", "validation", "test")]
|
||||
)
|
||||
history = load_jsonl(history_path)
|
||||
_validate_records(history, history_path)
|
||||
fixtures = load_classifiable_fixtures(fixtures_path)
|
||||
|
||||
eval_fixtures = [
|
||||
{"prompt": record["prompt"], "purpose": record["purpose"]}
|
||||
for split in ("validation", "test")
|
||||
for record in base[split]
|
||||
] + fixtures
|
||||
eval_hashes = {prompt_hash(record["prompt"]) for record in eval_fixtures}
|
||||
|
||||
eligible: list[SourceRecord] = []
|
||||
exclusions: list[dict[str, Any]] = []
|
||||
for record in _indexed_records(history, history_path):
|
||||
if record.value["slice"] == "vague-eval":
|
||||
exclusions.append(_exclusion(record, reason="vague-eval"))
|
||||
continue
|
||||
|
||||
key = normalized_key(record.value["prompt"])
|
||||
previous = base_keys.get(key)
|
||||
if previous is not None:
|
||||
split_name, previous_label = previous
|
||||
if previous_label != record.value["purpose"]:
|
||||
raise DataError(
|
||||
f"{history_path}:{record.line}: label conflicts with exact "
|
||||
f"{split_name} prompt"
|
||||
)
|
||||
exclusions.append(
|
||||
_exclusion(
|
||||
record,
|
||||
reason=(
|
||||
"exact-base-train-overlap"
|
||||
if split_name == "train"
|
||||
else "exact-eval-overlap"
|
||||
),
|
||||
matched_prompt_hash=prompt_hash(record.value["prompt"]),
|
||||
similarity=1.0,
|
||||
)
|
||||
)
|
||||
continue
|
||||
eligible.append(record)
|
||||
|
||||
curated = curate_records(
|
||||
eligible,
|
||||
eval_fixtures,
|
||||
near_duplicate_threshold=near_duplicate_threshold,
|
||||
)
|
||||
for duplicate in curated.duplicates:
|
||||
exclusions.append(
|
||||
_exclusion(
|
||||
duplicate.dropped,
|
||||
reason=(
|
||||
f"{duplicate.kind}-eval-overlap"
|
||||
if duplicate.matched_prompt_hash in eval_hashes
|
||||
else f"{duplicate.kind}-history-duplicate"
|
||||
),
|
||||
matched_prompt_hash=duplicate.matched_prompt_hash,
|
||||
similarity=duplicate.similarity,
|
||||
)
|
||||
)
|
||||
|
||||
accepted_history = [record.value for record in curated.records]
|
||||
output_train = base["train"] + accepted_history
|
||||
train_bytes = jsonl_bytes(output_train)
|
||||
validation_bytes = base_paths["validation"].read_bytes()
|
||||
test_bytes = base_paths["test"].read_bytes()
|
||||
|
||||
if output_dir.exists() and any(output_dir.iterdir()):
|
||||
if not overwrite_output:
|
||||
raise DataError(
|
||||
f"{output_dir}: output is not empty; pass --overwrite-output "
|
||||
"intentionally"
|
||||
)
|
||||
shutil.rmtree(output_dir)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
(output_dir / "train.jsonl").write_bytes(train_bytes)
|
||||
(output_dir / "validation.jsonl").write_bytes(validation_bytes)
|
||||
(output_dir / "test.jsonl").write_bytes(test_bytes)
|
||||
write_jsonl(
|
||||
output_dir / "history-exclusions.jsonl",
|
||||
sorted(exclusions, key=lambda value: value["sourceLine"]),
|
||||
)
|
||||
|
||||
exclusion_counts = dict(
|
||||
sorted(Counter(value["reason"] for value in exclusions).items())
|
||||
)
|
||||
accepted_sources = _indexed_records(accepted_history, history_path)
|
||||
manifest = {
|
||||
"schemaVersion": 1,
|
||||
"experimentVersion": EXPERIMENT_VERSION,
|
||||
"policy": {
|
||||
"historyUsage": "training-only",
|
||||
"vagueEval": "excluded from optimization",
|
||||
"exactBaseTrainOverlap": "excluded",
|
||||
"exactEvaluationOverlap": "excluded",
|
||||
"nearEvaluationOverlap": "excluded",
|
||||
"nearHistoryDuplicates": "excluded",
|
||||
"nearDuplicateThreshold": near_duplicate_threshold,
|
||||
"validationAndTestBytes": "identical to base dataset",
|
||||
},
|
||||
"sources": {
|
||||
"baseDataset": {
|
||||
"path": _relative(base_dataset),
|
||||
"splits": {
|
||||
split: {
|
||||
"path": _relative(path),
|
||||
"records": len(base[split]),
|
||||
"sha256": file_sha256(path),
|
||||
}
|
||||
for split, path in base_paths.items()
|
||||
},
|
||||
},
|
||||
"history": {
|
||||
"path": _relative(history_path),
|
||||
"records": len(history),
|
||||
"sha256": file_sha256(history_path),
|
||||
},
|
||||
"fixtures": {
|
||||
"path": _relative(fixtures_path),
|
||||
"classifiableRecords": len(fixtures),
|
||||
"sha256": file_sha256(fixtures_path),
|
||||
},
|
||||
},
|
||||
"augmentation": {
|
||||
"inputHistoryRecords": len(history),
|
||||
"acceptedHistoryRecords": len(accepted_history),
|
||||
"excludedHistoryRecords": len(exclusions),
|
||||
"exclusions": exclusion_counts,
|
||||
"acceptedDistribution": distribution(accepted_sources),
|
||||
},
|
||||
"outputs": {
|
||||
"train": {
|
||||
"records": len(output_train),
|
||||
"sha256": _sha256_bytes(train_bytes),
|
||||
},
|
||||
"validation": {
|
||||
"records": len(base["validation"]),
|
||||
"sha256": _sha256_bytes(validation_bytes),
|
||||
},
|
||||
"test": {
|
||||
"records": len(base["test"]),
|
||||
"sha256": _sha256_bytes(test_bytes),
|
||||
},
|
||||
},
|
||||
}
|
||||
write_json(output_dir / "manifest.json", manifest)
|
||||
return manifest
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--base-dataset", type=Path, default=DEFAULT_BASE_DATASET)
|
||||
parser.add_argument("--history", type=Path, default=DEFAULT_HISTORY)
|
||||
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
|
||||
parser.add_argument("--near-duplicate-threshold", type=float, default=0.92)
|
||||
parser.add_argument("--overwrite-output", action="store_true")
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
try:
|
||||
manifest = prepare(
|
||||
base_dataset=args.base_dataset.resolve(),
|
||||
history_path=args.history.resolve(),
|
||||
fixtures_path=args.fixtures.resolve(),
|
||||
output_dir=args.output_dir.resolve(),
|
||||
near_duplicate_threshold=args.near_duplicate_threshold,
|
||||
overwrite_output=args.overwrite_output,
|
||||
)
|
||||
except (DataError, OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
augmentation = manifest["augmentation"]
|
||||
outputs = manifest["outputs"]
|
||||
print(
|
||||
f"Prepared {outputs['train']['records']} training records "
|
||||
f"({augmentation['acceptedHistoryRecords']} from history); validation and "
|
||||
f"test remain byte-identical to the base dataset."
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user