Merge nucleic/lucid-north-quail-rnvt into dev
This commit is contained in:
@@ -0,0 +1,77 @@
|
||||
import json
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(MODULE_DIR))
|
||||
|
||||
import prepare_data
|
||||
import purpose_data
|
||||
|
||||
|
||||
def example(index: int):
|
||||
return {
|
||||
"prompt": f"Implement sample endpoint number {index} with stable pagination",
|
||||
"purpose": purpose_data.LABELS[index % len(purpose_data.LABELS)],
|
||||
"secondary": None,
|
||||
"mixed": False,
|
||||
"difficulty": 0.4,
|
||||
"slice": "vague-eval" if index < 5 else "core",
|
||||
"lang": "en",
|
||||
}
|
||||
|
||||
|
||||
class PrepareIntegrationTests(unittest.TestCase):
|
||||
def test_refresh_then_verify_frozen_split(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
source = root / "source.jsonl"
|
||||
purpose_data.write_jsonl(source, (example(index) for index in range(80)))
|
||||
fixtures = root / "fixtures.json"
|
||||
fixtures.write_text(
|
||||
json.dumps(
|
||||
[
|
||||
{"prompt": "Plan the cache migration", "purpose": "planning"},
|
||||
{"prompt": "Anything else?", "purpose": "general"},
|
||||
]
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
output = root / "output"
|
||||
frozen = root / "frozen.jsonl"
|
||||
manifest = root / "manifest.json"
|
||||
|
||||
first = prepare_data.prepare(
|
||||
sources=[source],
|
||||
fixtures_path=fixtures,
|
||||
output_dir=output,
|
||||
frozen_test_path=frozen,
|
||||
manifest_path=manifest,
|
||||
refresh_frozen_test=True,
|
||||
seed=23,
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
second = prepare_data.prepare(
|
||||
sources=[source],
|
||||
fixtures_path=fixtures,
|
||||
output_dir=output,
|
||||
frozen_test_path=frozen,
|
||||
manifest_path=manifest,
|
||||
refresh_frozen_test=False,
|
||||
seed=23,
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
|
||||
self.assertEqual(first, second)
|
||||
self.assertEqual(65, first["splits"]["train"]["records"])
|
||||
self.assertEqual(8, first["splits"]["validation"]["records"])
|
||||
self.assertEqual(8, first["splits"]["test"]["logicalRecords"])
|
||||
train = purpose_data.load_jsonl(output / "train.jsonl")
|
||||
self.assertFalse(any(row["slice"] == "vague-eval" for row in train))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user