Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user