Files
nucleic-purpose-classifier/tests/test_audit_data.py
T

93 lines
2.8 KiB
Python
Raw Normal View History

import sys
import tempfile
import unittest
from pathlib import Path
import numpy as np
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import audit_data
from purpose_data import SourceRecord
def record(index, purpose, slice_name="core", language="en"):
return SourceRecord(
value={
"prompt": f"prompt {index}",
"purpose": purpose,
"secondary": None,
"mixed": False,
"difficulty": 0.5,
"slice": slice_name,
"lang": language,
},
source=Path("source.jsonl"),
line=index + 1,
)
class ReviewSampleTests(unittest.TestCase):
def test_sample_is_exact_and_deterministic(self):
records = [
record(index, audit_data.LABELS[index % len(audit_data.LABELS)])
for index in range(101)
]
first = audit_data.stratified_review_sample(records, fraction=0.1, seed=42)
second = audit_data.stratified_review_sample(records, fraction=0.1, seed=42)
self.assertEqual(10, len(first))
self.assertEqual(
[item.value["prompt"] for item in first],
[item.value["prompt"] for item in second],
)
def test_review_csv_has_blank_reviewer_fields(self):
records = [record(0, "planning")]
with tempfile.TemporaryDirectory() as temporary:
path = Path(temporary) / "review.csv"
audit_data.write_review_csv(path, records)
text = path.read_text(encoding="utf-8")
self.assertIn("reviewedPurpose", text)
self.assertIn("prompt 0", text)
class SemanticCandidateTests(unittest.TestCase):
def test_threshold_depends_on_label_agreement(self):
embeddings = np.asarray(
[
[1.0, 0.0],
[0.98, 0.2],
[0.98, -0.2],
],
dtype=np.float32,
)
candidates = audit_data.semantic_candidates(
embeddings,
["planning", "planning", "writing"],
same_label_threshold=0.97,
cross_label_threshold=0.99,
neighbors=2,
block_size=2,
)
pairs = {(item.left, item.right) for item in candidates}
self.assertIn((0, 1), pairs)
self.assertNotIn((0, 2), pairs)
def test_candidate_pairs_are_deduplicated(self):
embeddings = np.asarray([[1.0, 0.0], [1.0, 0.0]], dtype=np.float32)
candidates = audit_data.semantic_candidates(
embeddings,
["review", "review"],
same_label_threshold=0.9,
cross_label_threshold=0.9,
neighbors=1,
block_size=1,
)
self.assertEqual(1, len(candidates))
if __name__ == "__main__":
unittest.main()