Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 23 additions & 2 deletions docs/source/content/sparse_probing.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -120,15 +132,24 @@ 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
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.

Expand Down
102 changes: 102 additions & 0 deletions tests/unit/tools/test_sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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():
Expand Down Expand Up @@ -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
Expand Down
50 changes: 49 additions & 1 deletion transformer_lens/tools/analysis/sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -38,6 +45,8 @@ class SparseProbeMetrics:
precision: float
recall: float
f1: float
roc_auc: float
average_precision: float


@dataclass(frozen=True)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand All @@ -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),
)


Expand Down Expand Up @@ -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"),
)


Expand Down
Loading