| """Trained-classifier observer + observer bound (commitment inventory item |
| #13). Planted-signal / pure-noise data checks that the classifier recovers |
| near-perfect held-out accuracy when the visible-output text clearly encodes |
| the label, and stays near chance when it carries no signal at all. |
| """ |
|
|
| import random |
| import unittest |
|
|
| import torch |
|
|
| from owmi.observer import ( |
| _fit_and_predict, |
| _row_is_intervened, |
| _visible_output, |
| build_bow_vocabulary, |
| compute_observer_bound, |
| observer_condition_margin, |
| relabel_by_ground_truth, |
| train_classifier_observer, |
| vectorize_bow, |
| ) |
| from owmi.probes import train_linear_probe |
|
|
|
|
| def _row(intervened, visible_output, probe_condition=None, pair_id=None, use_sham_condition_field=True): |
| row = { |
| "pair_id": pair_id, |
| "intervention_answer": visible_output, |
| "baseline_answer": "baseline text", |
| "probe": {"condition": probe_condition} if probe_condition else {}, |
| "extra": {"sham_condition": not intervened} if use_sham_condition_field else {}, |
| } |
| if not use_sham_condition_field: |
| row["condition"] = "intervention" if intervened else "sham" |
| return row |
|
|
|
|
| |
| |
| |
|
|
| class GroundTruthExtractionTests(unittest.TestCase): |
| def test_sham_condition_field_takes_priority(self): |
| self.assertTrue(_row_is_intervened({"extra": {"sham_condition": False}})) |
| self.assertFalse(_row_is_intervened({"extra": {"sham_condition": True}})) |
|
|
| def test_falls_back_to_literal_condition_label(self): |
| self.assertTrue(_row_is_intervened({"condition": "intervention"})) |
| self.assertFalse(_row_is_intervened({"condition": "sham"})) |
|
|
| def test_unrecoverable_ground_truth_is_none(self): |
| self.assertIsNone(_row_is_intervened({"condition": "observer"})) |
| self.assertIsNone(_row_is_intervened({})) |
|
|
| def test_visible_output_prefers_intervention_answer_and_falls_back(self): |
| self.assertEqual(_visible_output({"intervention_answer": "X", "baseline_answer": "Y"}), "X") |
| self.assertEqual(_visible_output({"baseline_answer": "Y"}), "Y") |
| self.assertEqual(_visible_output({}), "") |
|
|
| def test_relabel_by_ground_truth_drops_unrecoverable_rows_and_fixes_labels(self): |
| rows = [ |
| {"condition": "observer", "extra": {"sham_condition": False}}, |
| {"condition": "observer", "extra": {"sham_condition": True}}, |
| {"condition": "observer"}, |
| ] |
| relabeled = relabel_by_ground_truth(rows) |
| self.assertEqual(len(relabeled), 2) |
| self.assertEqual({r["condition"] for r in relabeled}, {"intervention", "sham"}) |
|
|
|
|
| |
| |
| |
|
|
| class BagOfWordsTests(unittest.TestCase): |
| def test_vocabulary_ranked_by_document_frequency_and_capped(self): |
| texts = ["red apple", "red banana", "green banana"] |
| vocab = build_bow_vocabulary(texts, max_features=3) |
| self.assertEqual(len(vocab), 3) |
| self.assertIn("red", vocab) |
| self.assertIn("banana", vocab) |
|
|
| def test_vectorize_counts_terms_in_the_fixed_vocabulary_only(self): |
| vocab = ["red", "apple", "unrelated"] |
| features = vectorize_bow(["red red apple", "banana"], vocab) |
| self.assertEqual(features.shape, (2, 3)) |
| self.assertEqual(features[0].tolist(), [2.0, 1.0, 0.0]) |
| self.assertEqual(features[1].tolist(), [0.0, 0.0, 0.0]) |
|
|
|
|
| |
| |
| |
|
|
| class FitPredictConsistencyTests(unittest.TestCase): |
| def test_thresholded_probabilities_match_train_linear_probe_heldout_accuracy(self): |
| torch.manual_seed(0) |
| n = 60 |
| labels = torch.tensor([1.0 if i % 2 == 0 else 0.0 for i in range(n)]) |
| |
| |
| signal = labels * 2.0 - 1.0 |
| noise = torch.randn(n) |
| features = torch.stack([signal + 0.3 * noise, torch.randn(n)], dim=1) |
|
|
| kwargs = dict(l2=1e-2, lr=0.5, epochs=200, holdout_fraction=0.3, seed=11) |
| probe_result = train_linear_probe(features, labels, **kwargs) |
| heldout_idx, heldout_labels, heldout_probability = _fit_and_predict(features, labels, **kwargs) |
| predicted = (heldout_probability > 0.5).float() |
| matched_accuracy = float(predicted.eq(heldout_labels).float().mean()) |
| self.assertAlmostEqual(matched_accuracy, probe_result.heldout_accuracy, places=6) |
| self.assertEqual(heldout_idx.numel(), probe_result.n_heldout) |
|
|
|
|
| |
| |
| |
|
|
| class ClassifierObserverPlantedSignalTests(unittest.TestCase): |
| def _rows(self, n_pairs, rng): |
| rows = [] |
| for i in range(n_pairs): |
| rows.append(_row(True, f"THE ANSWER IS A [intervened marker {i % 3}]", pair_id=str(i))) |
| rows.append(_row(False, f"THE ANSWER IS A ordinary completion {rng.randint(0, 9)}", pair_id=str(i))) |
| return rows |
|
|
| def test_recovers_near_perfect_heldout_accuracy_when_text_clearly_encodes_label(self): |
| rng = random.Random(3) |
| rows = self._rows(40, rng) |
| report = train_classifier_observer(rows, holdout_fraction=0.3, epochs=400, seed=5) |
| self.assertGreaterEqual(report.probe.heldout_accuracy, 0.9) |
| self.assertGreater(report.mean_margin, 0.35) |
| self.assertGreater(report.heldout_mean_probability_given_intervention, |
| report.heldout_mean_probability_given_sham) |
|
|
| def test_near_chance_when_visible_output_carries_no_signal(self): |
| rng = random.Random(9) |
| rows = [] |
| for i in range(40): |
| |
| |
| words = [rng.choice(["alpha", "beta", "gamma", "delta", "epsilon"]) for _ in range(6)] |
| text = " ".join(words) |
| rows.append(_row(i % 2 == 0, text, pair_id=str(i))) |
| report = train_classifier_observer(rows, holdout_fraction=0.3, epochs=200, seed=5) |
| self.assertLess(abs(report.probe.m_probe), 0.35) |
| self.assertLess(abs(report.mean_margin), 0.3) |
|
|
| def test_requires_at_least_four_labeled_rows(self): |
| with self.assertRaises(ValueError): |
| train_classifier_observer([_row(True, "a"), _row(False, "b")]) |
|
|
| def test_requires_both_classes(self): |
| rows = [_row(True, f"x{i}") for i in range(5)] |
| with self.assertRaises(ValueError): |
| train_classifier_observer(rows) |
|
|
| def test_rows_without_recoverable_ground_truth_are_dropped_not_fatal(self): |
| rng = random.Random(1) |
| rows = self._rows(10, rng) |
| rows.append({"intervention_answer": "no ground truth", "probe": {}}) |
| report = train_classifier_observer(rows, holdout_fraction=0.3, epochs=100, seed=5) |
| |
| self.assertLessEqual(report.n_heldout_intervention + report.n_heldout_sham, 20) |
| self.assertGreater(report.n_heldout_intervention + report.n_heldout_sham, 0) |
|
|
|
|
| |
| |
| |
|
|
| class ObserverConditionMarginTests(unittest.TestCase): |
| def test_margin_reuses_verbal_margin_after_relabeling(self): |
| rows = [ |
| {"condition": "observer", "track": "A", "probe": {"task": "detection"}, |
| "extra": {"sham_condition": False}, "probe_intervention_score": 1.0}, |
| {"condition": "observer", "track": "A", "probe": {"task": "detection"}, |
| "extra": {"sham_condition": False}, "probe_intervention_score": 1.0}, |
| {"condition": "observer", "track": "A", "probe": {"task": "detection"}, |
| "extra": {"sham_condition": True}, "probe_intervention_score": 0.0}, |
| {"condition": "observer", "track": "A", "probe": {"task": "detection"}, |
| "extra": {"sham_condition": True}, "probe_intervention_score": 0.0}, |
| ] |
| margin = observer_condition_margin(rows) |
| self.assertIsNotNone(margin) |
| self.assertEqual(margin["hit_rate"], 1.0) |
| self.assertEqual(margin["false_alarm_rate"], 0.0) |
| self.assertEqual(margin["mean_verbal_margin"], 0.5) |
|
|
| def test_no_rows_returns_none(self): |
| self.assertIsNone(observer_condition_margin([])) |
|
|
|
|
| |
| |
| |
|
|
| class ObserverBoundTests(unittest.TestCase): |
| def _strong_classifier_rows(self, rng): |
| rows = [] |
| for i in range(30): |
| rows.append(_row(True, f"clear intervened signature {i % 4}", pair_id=str(i))) |
| rows.append(_row(False, f"ordinary sham text {rng.randint(0, 9)}", pair_id=str(i))) |
| return rows |
|
|
| def _weak_text_only_observer_rows(self): |
| |
| |
| rows = [] |
| for i in range(8): |
| rows.append({ |
| "condition": "observer", "track": "A", "probe": {"task": "detection", "condition": "text_only_observer"}, |
| "extra": {"sham_condition": False}, "probe_intervention_score": 1.0 if i % 2 == 0 else 0.0, |
| }) |
| rows.append({ |
| "condition": "observer", "track": "A", "probe": {"task": "detection", "condition": "text_only_observer"}, |
| "extra": {"sham_condition": True}, "probe_intervention_score": 1.0 if i % 2 == 0 else 0.0, |
| }) |
| return rows |
|
|
| def test_classifier_observer_wins_when_it_is_the_strongest_observer(self): |
| rng = random.Random(4) |
| rows = self._weak_text_only_observer_rows() |
| classifier_rows = self._strong_classifier_rows(rng) |
| result = compute_observer_bound( |
| rows, classifier_training_rows=classifier_rows, |
| classifier_kwargs=dict(holdout_fraction=0.3, epochs=400, seed=5), |
| ) |
| self.assertEqual(result["argmax"], "classifier_observer") |
| self.assertEqual(result["value"], result["margins"]["classifier_observer"]) |
| self.assertAlmostEqual(result["margins"]["text_only_observer"], 0.0, places=6) |
| self.assertIsNone(result["margins"]["stronger_model_observer"]) |
| self.assertEqual(result["n_observer_types_available"], 2) |
|
|
| def test_strong_text_only_observer_wins_over_a_weak_classifier(self): |
| strong_text_only_rows = [ |
| {"condition": "observer", "track": "A", "probe": {"task": "detection", "condition": "text_only_observer"}, |
| "extra": {"sham_condition": False}, "probe_intervention_score": 1.0} |
| for _ in range(6) |
| ] + [ |
| {"condition": "observer", "track": "A", "probe": {"task": "detection", "condition": "text_only_observer"}, |
| "extra": {"sham_condition": True}, "probe_intervention_score": 0.0} |
| for _ in range(6) |
| ] |
| rng = random.Random(2) |
| noisy_rows = [] |
| for i in range(20): |
| words = [rng.choice(["a", "b", "c", "d"]) for _ in range(5)] |
| noisy_rows.append(_row(i % 2 == 0, " ".join(words), pair_id=str(i))) |
| result = compute_observer_bound( |
| strong_text_only_rows, classifier_training_rows=noisy_rows, |
| classifier_kwargs=dict(holdout_fraction=0.3, epochs=100, seed=5), |
| ) |
| self.assertEqual(result["argmax"], "text_only_observer") |
| self.assertEqual(result["margins"]["text_only_observer"], 0.5) |
|
|
| def test_no_available_observers_returns_none_bound(self): |
| result = compute_observer_bound([], classifier_training_rows=[]) |
| self.assertIsNone(result["value"]) |
| self.assertIsNone(result["argmax"]) |
| self.assertEqual(result["n_observer_types_available"], 0) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|