diff --git a/docs/source/content/sparse_probing.md b/docs/source/content/sparse_probing.md index b8c5d2832..f348a5103 100644 --- a/docs/source/content/sparse_probing.md +++ b/docs/source/content/sparse_probing.md @@ -83,6 +83,18 @@ $\alpha_c=1$. The intercept is not regularized. Positive predictions have nonneg Accuracy, precision, recall, F1, and all four confusion counts are returned; precision or F1 is zero when its denominator is zero. F1 is the primary sparse-probing metric. +F1 and the other threshold metrics describe a single operating point, `logit >= 0`, so they can +reward a probe that carries no information. When every selected coordinate is constant on the +training rows, the fit returns zero coefficients and a zero intercept, every held-out logit is +exactly `0.0`, and every held-out example is predicted positive. F1 then equals $2p/(1+p)$ for +the held-out positive rate $p$ (0.667 at $p = 0.5$) even though the probe cannot rank examples. + +`roc_auc` and `average_precision` are also returned, computed from the held-out logits without a +threshold. Both are tie-aware: tied logits share their average rank for ROC-AUC and form one +threshold for average precision. The degenerate probe above therefore scores ROC-AUC 0.5 and +average precision $p$, the chance levels. Check `constant_features` and compare against these +threshold-free metrics before reading a high F1 as decodability. + Feature-score reductions use float64 for float64 inputs and float32 otherwise. Selected matrices move to CPU float64, where LBFGS (at most `max_iter` iterations, with a fixed internal gradient stop) is followed by up to `max_refinement_steps` damped Newton steps on the `(k+1)`-square @@ -120,7 +132,13 @@ for k, probe, random_control in zip( sweep.random_coordinate_controls, strict=True, ): - print(k, probe.metrics.f1, random_control.f1.median()) + print( + k, + probe.metrics.f1, + probe.metrics.roc_auc, + random_control.f1.median(), + random_control.roc_auc.median(), + ) ``` Every k uses the same split, preprocessing mode, and L2 strength. `ks` must be strictly @@ -128,7 +146,10 @@ increasing and unique. Random-coordinate controls sample k distinct coordinates and fit the same classifier. Label-shuffle controls permute training labels, repeat selection and fitting, and evaluate against -the untouched held-out labels. The API returns raw control supports and metric distributions; it +the untouched held-out labels. Each control carries per-repeat accuracy, precision, recall, F1, +ROC-AUC, and average precision. A random-coordinate control that lands on a dead coordinate keeps +the inflated F1 described above, so compare controls on ROC-AUC as well. The API returns raw +control supports and metric distributions; it does not convert them into p-values or representation labels. A repeat count of zero disables that control. diff --git a/tests/unit/tools/test_sparse_probing.py b/tests/unit/tools/test_sparse_probing.py index a27254b34..30c350819 100644 --- a/tests/unit/tools/test_sparse_probing.py +++ b/tests/unit/tools/test_sparse_probing.py @@ -505,6 +505,95 @@ def test_binary_metrics_zero_division_policy(): assert metrics.precision == 0 assert metrics.recall == 0 assert metrics.f1 == 0 + assert metrics.roc_auc == 1.0 + assert metrics.average_precision == 1.0 + + +def test_binary_metrics_all_negative_labels_leave_both_threshold_free_scores_undefined(): + metrics = _binary_metrics(torch.tensor([-1.0, 0.5, 2.0]), torch.zeros(3, dtype=torch.int64)) + + assert math.isnan(metrics.roc_auc) + assert math.isnan(metrics.average_precision) + + +def test_binary_metrics_all_positive_labels_give_nan_roc_auc_and_unit_average_precision(): + # ROC-AUC needs a negative to rank against; average precision is still defined, since + # every threshold has precision one. + metrics = _binary_metrics(torch.tensor([-1.0, 0.5, 2.0]), torch.ones(3, dtype=torch.int64)) + + assert math.isnan(metrics.roc_auc) + assert metrics.average_precision == 1.0 + + +def test_dead_coordinate_probe_scores_chance_on_threshold_free_metrics(): + # Fixture from issue #1814: a constant selected coordinate fits to all-zero logits, + # which predicts every held-out example positive. + features = torch.randn(300, 16, generator=torch.Generator().manual_seed(0)) + features[:, 0] = 0.0 + labels = (torch.rand(300, generator=torch.Generator().manual_seed(0)) < 0.5).long() + + result = fit_sparse_probe(features[:, :1], labels, k=1, seed=0) + + assert result.constant_features.tolist() == [True] + assert float(result.coefficients[0]) == 0.0 + assert float(result.intercept) == 0.0 + metrics = result.metrics + positive_rate = result.test_positive_count / ( + result.test_positive_count + result.test_negative_count + ) + assert metrics.false_negatives == metrics.true_negatives == 0 + assert metrics.f1 == pytest.approx(2 * positive_rate / (1 + positive_rate)) + assert metrics.f1 == pytest.approx(0.7, abs=5e-4) + assert metrics.roc_auc == 0.5 + assert metrics.average_precision == pytest.approx(positive_rate) + + +@pytest.mark.parametrize( + ("logits", "labels", "roc_auc", "average_precision"), + [ + ([3.0, 2.0, 1.0, 0.0], [1, 1, 0, 0], 1.0, 1.0), + ([0.0, 1.0, 2.0, 3.0], [1, 1, 0, 0], 0.0, (1 / 3 + 2 / 4) / 2), + ([1.0, 1.0, 1.0, 1.0], [1, 0, 1, 0], 0.5, 0.5), + # The top two logits tie across classes: they are one threshold at precision 1/2, + # and one positive-negative pair counts half. + ([2.0, 2.0, 1.0, 0.0], [1, 0, 1, 0], 0.625, (1 / 2 + 2 / 3) / 2), + ], +) +def test_threshold_free_metrics_match_hand_computed_values( + logits, labels, roc_auc, average_precision +): + metrics = _binary_metrics(torch.tensor(logits, dtype=torch.float64), torch.tensor(labels)) + + assert metrics.roc_auc == pytest.approx(roc_auc) + assert metrics.average_precision == pytest.approx(average_precision) + + +@pytest.mark.parametrize("seed", range(5)) +def test_threshold_free_metrics_match_a_brute_force_tie_aware_reference(seed): + generator = torch.Generator().manual_seed(seed) + labels = torch.arange(60) % 3 == 0 + # Coarse rounding makes many ties, both within and across classes. + logits = torch.round(torch.randn(60, generator=generator, dtype=torch.float64) + labels) + + positive_logits = logits[labels] + negative_logits = logits[~labels] + pair_wins = (positive_logits[:, None] > negative_logits[None, :]).double() + pair_ties = (positive_logits[:, None] == negative_logits[None, :]).double() + expected_roc_auc = float((pair_wins + 0.5 * pair_ties).mean()) + + expected_average_precision = 0.0 + previous_recall = 0.0 + for threshold in sorted(set(logits.tolist()), reverse=True): + predicted = logits >= threshold + precision = float((predicted & labels).sum()) / float(predicted.sum()) + recall = float((predicted & labels).sum()) / float(labels.sum()) + expected_average_precision += (recall - previous_recall) * precision + previous_recall = recall + + metrics = _binary_metrics(logits, labels) + + assert metrics.roc_auc == pytest.approx(expected_roc_auc, abs=1e-12) + assert metrics.average_precision == pytest.approx(expected_average_precision, abs=1e-12) @pytest.mark.parametrize( @@ -621,10 +710,14 @@ def test_disabled_controls_return_empty_aligned_results(): random_control.precision, random_control.recall, random_control.f1, + random_control.roc_auc, + random_control.average_precision, shuffle_control.accuracy, shuffle_control.precision, shuffle_control.recall, shuffle_control.f1, + shuffle_control.roc_auc, + shuffle_control.average_precision, ): assert metric_values.shape == (0,) assert metric_values.dtype == torch.float64 @@ -662,6 +755,9 @@ def test_controls_are_deterministic_use_unique_supports_and_do_not_touch_global_ assert torch.equal(left.precision, right.precision) assert torch.equal(left.recall, right.recall) assert torch.equal(left.f1, right.f1) + assert torch.equal(left.roc_auc, right.roc_auc) + assert torch.equal(left.average_precision, right.average_precision) + assert left.roc_auc.shape == left.average_precision.shape == (4,) for support in left.supports: assert torch.unique(support).numel() == 2 @@ -682,6 +778,10 @@ def test_controls_remain_below_a_strong_planted_feature(): assert actual_f1 > 0.98 assert float(sweep.random_coordinate_controls[0].f1.median()) < actual_f1 - 0.2 assert float(sweep.label_shuffle_controls[0].f1.median()) < actual_f1 - 0.2 + actual_roc_auc = sweep.results[0].metrics.roc_auc + assert actual_roc_auc > 0.98 + assert float(sweep.random_coordinate_controls[0].roc_auc.median()) < actual_roc_auc - 0.2 + assert float(sweep.label_shuffle_controls[0].roc_auc.median()) < actual_roc_auc - 0.2 def test_larger_k_improves_distributed_decodability_without_assigning_a_representation_label(): @@ -756,6 +856,8 @@ def test_label_shuffle_control_is_scored_against_the_true_heldout_labels(): assert float(control.precision[0]) == rescored.precision assert float(control.recall[0]) == rescored.recall assert float(control.f1[0]) == rescored.f1 + assert float(control.roc_auc[0]) == rescored.roc_auc + assert float(control.average_precision[0]) == rescored.average_precision # Scoring the same fit against a different held-out label alignment would move the # metrics, so the exact match above pins the scoring labels to the true held-out labels. assert _binary_metrics(logits, ~true_test_labels).accuracy != rescored.accuracy diff --git a/transformer_lens/tools/analysis/sparse_probing.py b/transformer_lens/tools/analysis/sparse_probing.py index 4aa7e9b67..d09c30917 100644 --- a/transformer_lens/tools/analysis/sparse_probing.py +++ b/transformer_lens/tools/analysis/sparse_probing.py @@ -28,7 +28,14 @@ @dataclass(frozen=True) class SparseProbeMetrics: - """Held-out binary-classification metrics and confusion counts.""" + """Held-out binary-classification metrics and confusion counts. + + Confusion counts, accuracy, precision, recall, and F1 are read at the logit-zero + threshold. ``roc_auc`` and ``average_precision`` are threshold-free and tie-aware: + held-out examples with equal logits share one rank and one threshold. The stratified + split always holds out both classes; if a class is absent, ``roc_auc`` is NaN, and + ``average_precision`` is NaN without positives and 1.0 without negatives. + """ true_positives: int true_negatives: int @@ -38,6 +45,8 @@ class SparseProbeMetrics: precision: float recall: float f1: float + roc_auc: float + average_precision: float @dataclass(frozen=True) @@ -91,6 +100,8 @@ class SparseProbeControl: precision: Float[torch.Tensor, "repeat"] recall: Float[torch.Tensor, "repeat"] f1: Float[torch.Tensor, "repeat"] + roc_auc: Float[torch.Tensor, "repeat"] + average_precision: Float[torch.Tensor, "repeat"] @dataclass(frozen=True) @@ -522,6 +533,39 @@ def closure() -> torch.Tensor: ) +def _roc_auc(logits: torch.Tensor, positive: torch.Tensor) -> float: + """Mann-Whitney ROC-AUC with average ranks for tied logits; NaN if a class is absent.""" + positive_count = int(positive.sum().item()) + negative_count = positive.numel() - positive_count + if positive_count == 0 or negative_count == 0: + return math.nan + _, inverse, counts = torch.unique(logits, return_inverse=True, return_counts=True) + counts = counts.to(dtype=torch.float64) + # A tie group occupying ranks start+1..end gets their mean, end - (count - 1) / 2. + average_ranks = (counts.cumsum(0) - (counts - 1) / 2)[inverse] + positive_rank_sum = float(average_ranks[positive].sum().item()) + return (positive_rank_sum - positive_count * (positive_count + 1) / 2) / ( + positive_count * negative_count + ) + + +def _average_precision(logits: torch.Tensor, positive: torch.Tensor) -> float: + """Step-wise average precision over distinct logits; NaN if there is no positive.""" + positive_count = int(positive.sum().item()) + if positive_count == 0: + return math.nan + _, inverse = torch.unique(logits, return_inverse=True) + group_count = int(inverse.max().item()) + 1 + # Tied logits form one threshold, so the recall gained there is scored at a single precision. + group_positives = torch.zeros(group_count, dtype=torch.float64).index_add_( + 0, inverse, positive.to(dtype=torch.float64) + ) + group_sizes = torch.bincount(inverse, minlength=group_count).to(dtype=torch.float64) + positives_descending = group_positives.flip(0) + precision = positives_descending.cumsum(0) / group_sizes.flip(0).cumsum(0) + return float((positives_descending * precision).sum().item()) / positive_count + + def _binary_metrics(logits: torch.Tensor, labels: torch.Tensor) -> SparseProbeMetrics: predictions = logits >= 0 positive = labels.to(dtype=torch.bool) @@ -546,6 +590,8 @@ def _binary_metrics(logits: torch.Tensor, labels: torch.Tensor) -> SparseProbeMe precision=precision, recall=recall, f1=f1, + roc_auc=_roc_auc(logits, positive), + average_precision=_average_precision(logits, positive), ) @@ -657,6 +703,8 @@ def metric_tensor(name: str) -> torch.Tensor: precision=metric_tensor("precision"), recall=metric_tensor("recall"), f1=metric_tensor("f1"), + roc_auc=metric_tensor("roc_auc"), + average_precision=metric_tensor("average_precision"), )