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
8 changes: 8 additions & 0 deletions docs/source/content/sparse_probing.md
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,14 @@ control supports and metric distributions; it
does not convert them into p-values or representation labels. A repeat count of zero disables that
control.

Unlike the main results, control distributions carry only raw metrics — no per-fit convergence
diagnostics. A control fit that fails to converge is not fatal to the sweep: it is excluded from
its control distribution and recorded in `sweep.rejections` (a `SparseProbeRejection` naming the
arm, `k`, repeat, support, and reason), so the completed main probes and other controls survive.
Each control's rows therefore cover only the repeats that converged. Main-probe fits are not made
partial this way — a main fit that misses the acceptance threshold still raises, as does
`fit_sparse_probe`.

Controls can be expensive: the sweep performs one main fit plus both requested control counts for
every k. Start with small grids and repeat counts.

Expand Down
48 changes: 48 additions & 0 deletions tests/unit/tools/test_sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -886,6 +886,54 @@ def test_sweep_rejects_invalid_grid_and_control_counts(ks, kwargs, message):
sweep_sparse_probe(features, labels, ks=ks, **kwargs)


def test_sweep_records_rejected_controls_and_keeps_completed_fits():
# max_refinement_steps=0 with a strict tolerance makes some control fits miss the
# acceptance threshold. Before the partial-result path a rejection raised and discarded
# the whole sweep — including the completed main probes, which are the expensive part.
features = torch.randn(200, 32, generator=torch.Generator().manual_seed(0))
labels = (torch.rand(200, generator=torch.Generator().manual_seed(1)) < 0.5).long()
features[:, 3] += labels.float() * 1.5
strict = dict(seed=0, max_refinement_steps=0, decrement_tolerance=1e-14)

sweep = sweep_sparse_probe(
features, labels, ks=[1, 2, 4], n_random_subsets=10, n_label_shuffles=10, **strict
)

# Main probes for every k survive.
assert tuple(result.k for result in sweep.results) == (1, 2, 4)

# At least one control was rejected and recorded with an actionable reason.
assert len(sweep.rejections) > 0
for rejection in sweep.rejections:
assert rejection.arm in ("random_coordinate", "label_shuffle")
assert rejection.k in (1, 2, 4)
assert "did not converge" in rejection.reason
assert rejection.support.numel() == rejection.k

# Each control distribution keeps exactly the repeats that converged: the requested
# count minus the rejections recorded for that arm and k.
for arm, controls in (
("random_coordinate", sweep.random_coordinate_controls),
("label_shuffle", sweep.label_shuffle_controls),
):
for k, control in zip(sweep.ks, controls, strict=True):
rejected = sum(1 for r in sweep.rejections if r.arm == arm and r.k == k)
assert control.supports.shape[0] == 10 - rejected
assert control.f1.numel() == 10 - rejected


def test_sweep_has_no_rejections_when_every_control_converges():
features, labels = _planted_data(n_examples=160, n_features=12)

sweep = sweep_sparse_probe(
features, labels, ks=[1, 2], n_random_subsets=4, n_label_shuffles=4, seed=5
)

assert sweep.rejections == ()
for control in (*sweep.random_coordinate_controls, *sweep.label_shuffle_controls):
assert control.supports.shape[0] == 4


def test_control_draws_depend_only_on_seed_k_and_arm():
# A k=2 control must be a property of k=2 at a fixed seed: independent of which
# other k values were requested and of the other arm's repeat count. With one
Expand Down
2 changes: 2 additions & 0 deletions transformer_lens/tools/analysis/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,7 @@
from transformer_lens.tools.analysis.sparse_probing import (
SparseProbeControl,
SparseProbeMetrics,
SparseProbeRejection,
SparseProbeResult,
SparseProbeSweep,
fit_sparse_probe,
Expand Down Expand Up @@ -153,6 +154,7 @@
"RankReportRow",
"SparseProbeControl",
"SparseProbeMetrics",
"SparseProbeRejection",
"SparseProbeResult",
"SparseProbeSweep",
"SubspaceBasis",
Expand Down
85 changes: 63 additions & 22 deletions transformer_lens/tools/analysis/sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,7 +94,14 @@ class SparseProbeResult:

@dataclass(frozen=True)
class SparseProbeControl:
"""Raw held-out metric distributions for one control at one sparsity."""
"""Raw held-out metric distributions for one control at one sparsity.

Rows are the control repeats that converged; a repeat whose fit was rejected
is excluded here (its coordinates and reason live in ``SparseProbeSweep.rejections``),
so the row count can be below the requested repeat count. These are raw metrics
only and carry no per-fit convergence diagnostics (``newton_decrement`` etc.) —
those live on ``SparseProbeResult`` for the main fits.
"""

supports: Int[torch.Tensor, "repeat selected_feature"]
accuracy: Float[torch.Tensor, "repeat"]
Expand All @@ -105,15 +112,38 @@ class SparseProbeControl:
average_precision: Float[torch.Tensor, "repeat"]


@dataclass(frozen=True)
class SparseProbeRejection:
"""A control fit excluded from the sweep because it did not converge.

The sweep records rejected control fits here instead of aborting, so completed
main probes and other controls survive. ``arm`` is ``"random_coordinate"`` or
``"label_shuffle"``, ``support`` is the coordinate set the rejected fit used, and
``reason`` is the convergence-failure message. Main-probe fits are not made
partial this way — they still raise (see ``sweep_sparse_probe``).
"""

arm: str
k: int
repeat: int
support: Int[torch.Tensor, "selected_feature"]
reason: str


@dataclass(frozen=True)
class SparseProbeSweep:
"""Probe results and aligned controls over a strictly increasing k-grid."""
"""Probe results and aligned controls over a strictly increasing k-grid.

``rejections`` lists control fits that did not converge and were excluded; it is
empty when every requested control fit converged.
"""

ks: tuple[int, ...]
results: tuple[SparseProbeResult, ...]
random_coordinate_controls: tuple[SparseProbeControl, ...]
label_shuffle_controls: tuple[SparseProbeControl, ...]
seed: int
rejections: tuple[SparseProbeRejection, ...] = ()


@dataclass(frozen=True)
Expand Down Expand Up @@ -850,11 +880,15 @@ def sweep_sparse_probe(
an estimate of the objective gap to the optimum in nats; a larger gap raises.

Returns:
Main probe results plus aligned raw control distributions.
Main probe results plus aligned raw control distributions. A control fit
that fails to converge is excluded and recorded in ``rejections`` rather
than aborting the sweep, so completed fits survive; each control's rows
cover only its converged repeats.

Raises:
ValueError: If the grid, controls, inputs, or options are invalid.
RuntimeError: If any main or control fit fails to converge.
RuntimeError: If a main-probe fit fails to converge. Control-fit failures
do not raise here — they are collected in ``SparseProbeSweep.rejections``.
"""
if isinstance(ks, (str, bytes)) or not isinstance(ks, Sequence):
raise ValueError("ks must be a non-empty sequence of positive integers")
Expand Down Expand Up @@ -902,23 +936,28 @@ def sweep_sparse_probe(

random_controls = []
shuffle_controls = []
rejections: list[SparseProbeRejection] = []
feature_count = validated.features.shape[1]
for k in k_values:
random_supports = []
random_metrics = []
for repeat in range(random_count):
draw_generator = _control_generator(validated.seed, k, "random", repeat)
support = torch.randperm(feature_count, generator=draw_generator)[:k].sort().values
random_supports.append(support)
random_metrics.append(
_fit_control(
validated,
train_indices,
test_indices,
support,
train_labels,
# A control fit that fails to converge is recorded and skipped, not fatal:
# losing one auxiliary draw must not discard the completed main probes and
# other controls. Main fits above are the primary output and still raise.
try:
metrics = _fit_control(
validated, train_indices, test_indices, support, train_labels
)
)
except RuntimeError as error:
rejections.append(
SparseProbeRejection("random_coordinate", k, repeat, support, str(error))
)
continue
random_supports.append(support)
random_metrics.append(metrics)
random_controls.append(_control_result(random_supports, random_metrics, k))

shuffle_supports = []
Expand All @@ -929,16 +968,17 @@ def sweep_sparse_probe(
shuffled_labels = train_labels[permutation]
shuffled_scores = _feature_scores(validated.features, shuffled_labels, train_indices)
support = torch.argsort(shuffled_scores.abs(), descending=True, stable=True)[:k]
shuffle_supports.append(support)
shuffle_metrics.append(
_fit_control(
validated,
train_indices,
test_indices,
support,
shuffled_labels,
try:
metrics = _fit_control(
validated, train_indices, test_indices, support, shuffled_labels
)
)
except RuntimeError as error:
rejections.append(
SparseProbeRejection("label_shuffle", k, repeat, support, str(error))
)
continue
shuffle_supports.append(support)
shuffle_metrics.append(metrics)
shuffle_controls.append(_control_result(shuffle_supports, shuffle_metrics, k))

return SparseProbeSweep(
Expand All @@ -947,4 +987,5 @@ def sweep_sparse_probe(
random_coordinate_controls=tuple(random_controls),
label_shuffle_controls=tuple(shuffle_controls),
seed=validated.seed,
rejections=tuple(rejections),
)
Loading