Merge nucleic/brisk-umber-viper-rg2p into dev
This commit is contained in:
@@ -0,0 +1,173 @@
|
||||
import importlib.util
|
||||
import json
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
MODULE_PATH = Path(__file__).resolve().parents[1] / "validate-data.py"
|
||||
SPEC = importlib.util.spec_from_file_location("validate_data", MODULE_PATH)
|
||||
assert SPEC is not None and SPEC.loader is not None
|
||||
validate_data = importlib.util.module_from_spec(SPEC)
|
||||
sys.modules[SPEC.name] = validate_data
|
||||
SPEC.loader.exec_module(validate_data)
|
||||
|
||||
|
||||
def example(**overrides):
|
||||
value = {
|
||||
"prompt": "Add a health-check endpoint",
|
||||
"purpose": "backendImpl",
|
||||
"secondary": None,
|
||||
"mixed": False,
|
||||
"difficulty": 0.4,
|
||||
"slice": "core",
|
||||
"lang": "en",
|
||||
}
|
||||
value.update(overrides)
|
||||
return value
|
||||
|
||||
|
||||
class RecordValidationTests(unittest.TestCase):
|
||||
def test_accepts_valid_record(self):
|
||||
validator = validate_data.Validator(
|
||||
fixture_path=None, process_checks=False
|
||||
)
|
||||
|
||||
self.assertTrue(validator._check_record(Path("batch.jsonl"), 1, example()))
|
||||
self.assertEqual([], validator.issues)
|
||||
|
||||
def test_rejects_mixed_field_inconsistencies(self):
|
||||
validator = validate_data.Validator(
|
||||
fixture_path=None, process_checks=False
|
||||
)
|
||||
|
||||
valid = validator._check_record(
|
||||
Path("batch.jsonl"),
|
||||
7,
|
||||
example(mixed=True, secondary=None, slice="core"),
|
||||
)
|
||||
|
||||
self.assertFalse(valid)
|
||||
messages = [issue.message for issue in validator.issues]
|
||||
self.assertIn("mixed=true requires a secondary purpose", messages)
|
||||
self.assertIn("the mixed slice and mixed field must agree", messages)
|
||||
|
||||
def test_rejects_unknown_fields_and_non_finite_difficulty(self):
|
||||
validator = validate_data.Validator(
|
||||
fixture_path=None, process_checks=False
|
||||
)
|
||||
|
||||
valid = validator._check_record(
|
||||
Path("batch.jsonl"),
|
||||
3,
|
||||
example(difficulty=float("nan"), surprise="value"),
|
||||
)
|
||||
|
||||
self.assertFalse(valid)
|
||||
messages = [issue.message for issue in validator.issues]
|
||||
self.assertIn("unexpected fields: surprise", messages)
|
||||
self.assertIn(
|
||||
"difficulty must be a finite number from 0.0 to 1.0", messages
|
||||
)
|
||||
|
||||
def test_bcp47_structure(self):
|
||||
for tag in ("en", "pt-BR", "zh-Hans-CN", "de-DE-1996", "x-project"):
|
||||
with self.subTest(tag=tag):
|
||||
self.assertTrue(validate_data.is_bcp47(tag))
|
||||
for tag in ("", "e", "en_US", "-en", "en-", "en-☃"):
|
||||
with self.subTest(tag=tag):
|
||||
self.assertFalse(validate_data.is_bcp47(tag))
|
||||
|
||||
|
||||
class DatasetValidationTests(unittest.TestCase):
|
||||
def test_valid_complete_batch_has_no_errors(self):
|
||||
slices = (
|
||||
["core"] * 110
|
||||
+ ["boundary"] * 40
|
||||
+ ["mixed"] * 20
|
||||
+ ["pasted-context"] * 20
|
||||
+ ["vague-eval"] * 10
|
||||
)
|
||||
languages = ["en"] * 190 + [
|
||||
"es",
|
||||
"de",
|
||||
"fr",
|
||||
"pt",
|
||||
"zh",
|
||||
"ja",
|
||||
"es-MX",
|
||||
"de-DE",
|
||||
"fr-CA",
|
||||
"pt-BR",
|
||||
]
|
||||
rows = []
|
||||
for index, (slice_name, lang) in enumerate(zip(slices, languages)):
|
||||
if index < 30:
|
||||
prompt = f"sample{index} task"
|
||||
elif index < 40 or slice_name == "pasted-context":
|
||||
prompt = f"sample{index} " + "context " * 110
|
||||
else:
|
||||
prompt = (
|
||||
f"sample{index} perform a realistic scoped change in the "
|
||||
"project with the listed constraints"
|
||||
)
|
||||
purpose = validate_data.PURPOSES[index % len(validate_data.PURPOSES)]
|
||||
is_mixed = slice_name == "mixed"
|
||||
secondary = (
|
||||
validate_data.PURPOSES[(index + 1) % len(validate_data.PURPOSES)]
|
||||
if is_mixed
|
||||
else None
|
||||
)
|
||||
rows.append(
|
||||
example(
|
||||
prompt=prompt,
|
||||
purpose=purpose,
|
||||
secondary=secondary,
|
||||
mixed=is_mixed,
|
||||
slice=slice_name,
|
||||
lang=lang,
|
||||
)
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
batch = root / "batch.jsonl"
|
||||
batch.write_text(
|
||||
"".join(json.dumps(row) + "\n" for row in rows), encoding="utf-8"
|
||||
)
|
||||
validator = validate_data.Validator(
|
||||
batch_size=200,
|
||||
expected_total=200,
|
||||
fixture_path=None,
|
||||
process_checks=True,
|
||||
)
|
||||
records = validator.validate([batch], [root])
|
||||
|
||||
self.assertEqual(200, len(records))
|
||||
self.assertEqual(
|
||||
[],
|
||||
[issue for issue in validator.issues if issue.severity == "error"],
|
||||
)
|
||||
|
||||
def test_reports_bad_json_and_normalized_duplicates(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
batch = Path(directory) / "batch.jsonl"
|
||||
batch.write_text(
|
||||
json.dumps(example(prompt="Fix spacing")) + "\n"
|
||||
+ "{not json}\n"
|
||||
+ json.dumps(example(prompt=" fix spacing ")) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
validator = validate_data.Validator(
|
||||
fixture_path=None, process_checks=False
|
||||
)
|
||||
validator.validate([batch], [batch])
|
||||
|
||||
messages = [issue.message for issue in validator.issues]
|
||||
self.assertTrue(any(message.startswith("invalid JSON:") for message in messages))
|
||||
self.assertTrue(any(message.startswith("duplicate prompt;") for message in messages))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user