diff --git a/docs/source/content/sparse_probing.md b/docs/source/content/sparse_probing.md index f348a5103..2f2f10b48 100644 --- a/docs/source/content/sparse_probing.md +++ b/docs/source/content/sparse_probing.md @@ -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. diff --git a/tests/unit/tools/test_sparse_probing.py b/tests/unit/tools/test_sparse_probing.py index 98e58147d..6e79f20af 100644 --- a/tests/unit/tools/test_sparse_probing.py +++ b/tests/unit/tools/test_sparse_probing.py @@ -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 diff --git a/transformer_lens/tools/analysis/__init__.py b/transformer_lens/tools/analysis/__init__.py index f2df06265..8206c382a 100644 --- a/transformer_lens/tools/analysis/__init__.py +++ b/transformer_lens/tools/analysis/__init__.py @@ -100,6 +100,7 @@ from transformer_lens.tools.analysis.sparse_probing import ( SparseProbeControl, SparseProbeMetrics, + SparseProbeRejection, SparseProbeResult, SparseProbeSweep, fit_sparse_probe, @@ -153,6 +154,7 @@ "RankReportRow", "SparseProbeControl", "SparseProbeMetrics", + "SparseProbeRejection", "SparseProbeResult", "SparseProbeSweep", "SubspaceBasis", diff --git a/transformer_lens/tools/analysis/sparse_probing.py b/transformer_lens/tools/analysis/sparse_probing.py index d42231bf1..61f9e8409 100644 --- a/transformer_lens/tools/analysis/sparse_probing.py +++ b/transformer_lens/tools/analysis/sparse_probing.py @@ -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"] @@ -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) @@ -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") @@ -902,6 +936,7 @@ def sweep_sparse_probe( random_controls = [] shuffle_controls = [] + rejections: list[SparseProbeRejection] = [] feature_count = validated.features.shape[1] for k in k_values: random_supports = [] @@ -909,16 +944,20 @@ def sweep_sparse_probe( 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 = [] @@ -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( @@ -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), )