52 lines
1.5 KiB
Python
52 lines
1.5 KiB
Python
import sys
|
|
import unittest
|
|
from collections import Counter
|
|
from pathlib import Path
|
|
|
|
|
|
MODULE_DIR = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(MODULE_DIR))
|
|
|
|
import export
|
|
|
|
|
|
class CalibrationSampleTests(unittest.TestCase):
|
|
def test_sample_is_exact_deterministic_and_stratified(self):
|
|
records = []
|
|
for index in range(100):
|
|
records.append(
|
|
{
|
|
"prompt": f"prompt {index}",
|
|
"purpose": "planning" if index < 80 else "writing",
|
|
"slice": "core" if index % 2 else "boundary",
|
|
"lang": "en" if index % 5 else "fr",
|
|
}
|
|
)
|
|
first = export.stratified_calibration_sample(records, 25, seed=42)
|
|
second = export.stratified_calibration_sample(records, 25, seed=42)
|
|
self.assertEqual(25, len(first))
|
|
self.assertEqual(
|
|
[item["prompt"] for item in first],
|
|
[item["prompt"] for item in second],
|
|
)
|
|
purposes = Counter(item["purpose"] for item in first)
|
|
self.assertEqual({"planning": 20, "writing": 5}, dict(purposes))
|
|
|
|
def test_sample_caps_at_population(self):
|
|
records = [
|
|
{
|
|
"prompt": "one",
|
|
"purpose": "planning",
|
|
"slice": "core",
|
|
"lang": "en",
|
|
}
|
|
]
|
|
self.assertEqual(
|
|
records,
|
|
export.stratified_calibration_sample(records, 10, seed=1),
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|