Skip to content

fix sparse probe partial results - #1828

Merged
jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
MdSadiqMd:sadiq/fix-sparse-probe-partial-results
Sep 28, 2026
Merged

jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
MdSadiqMd:sadiq/fix-sparse-probe-partial-results

Conversation

@MdSadiqMd

Copy link
Copy Markdown
Contributor

Description

A fit in the sparse-probing sweep runs LBFGS then damped Newton refinement and is accepted only when the Newton decrement is at most decrement_tolerance (the acceptance rule that landed in #1774). A fit that misses it raises, and sweep_sparse_probe had no partial-result path, so a single rejected fit discarded everything already computed — including the main probes, which are the expensive part and typically fine. On the sweep path most fits are controls, so this "lose everything" failure was overwhelmingly triggered by an auxiliary control repeat.

X = torch.randn(200, 32, generator=torch.Generator().manual_seed(0))
y = (torch.rand(200, generator=torch.Generator().manual_seed(1)) < 0.5).long()
X[:, 3] += y.float() * 1.5
kw = dict(seed=0, max_refinement_steps=0, decrement_tolerance=1e-14)

sweep_sparse_probe(X, y, ks=[1, 2, 4], n_random_subsets=0,  n_label_shuffles=0,  **kw)  # all converge
sweep_sparse_probe(X, y, ks=[1, 2, 4], n_random_subsets=10, n_label_shuffles=10, **kw)  # RuntimeError on a control

Since the refinement landed this no longer fires at shipped defaults, so it's robustness rather than routine loss — but the failure mode is still "lose everything" for a control that stopped short. Related: stop_reason, newton_decrement, and refinement_steps are on SparseProbeResult but not on controls, and there was no way to tell from a result which control fits were rejected or why.

This PR:

  • Adds a partial-result path. Each control fit in the sweep is wrapped so a convergence RuntimeError records the fit and is skipped rather than aborting the sweep. Completed main probes and other controls survive. Each control's rows now cover only the repeats that converged (the row count can be below the requested repeat count).
  • Reports which fits were rejected and why. New SparseProbeRejection(arm, k, repeat, support, reason) dataclass; SparseProbeSweep gains rejections: tuple[SparseProbeRejection, ...] = () (defaulted, so existing construction is unaffected). Exported from transformer_lens.tools.analysis.
  • Keeps fit_sparse_probe and main-probe fits raising. Only the sweep's control fits are made resilient. Main fits stay 1:1 with ks and still raise on non-convergence, and fit_sparse_probe is untouched — a hard failure on the primary output stays loud.
  • Documents the diagnostics decision. The proposal allowed either surfacing per-repeat diagnostics on SparseProbeControl or saying plainly that it doesn't; this takes the documentation route. SparseProbeControl's docstring and the sparse-probing guide now state that controls carry raw metrics only (no per-fit convergence diagnostics), and that rejected control fits are collected in sweep.rejections.

Acceptance, all verified:

  • A sweep with one (here, several) rejected control returns its completed fits — results present for every k, other controls intact.
  • New test covering the path (test_sweep_records_rejected_controls_and_keeps_completed_fits), which fails on the pre-fix code (the sweep raised) and passes now; plus a clean-path test asserting rejections == () when every control converges.
  • make unit-test passes (6695 passed, 56 skipped, 6 xfailed).
  • uv run mypy . passes (385 files).

Fixes #1815

Type of change

Please delete options that are not relevant.

  • Bug fix (non-breaking change which fixes an issue)
  • This change requires a documentation update

Screenshots

N/A

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

@jlarson4

Copy link
Copy Markdown
Collaborator

@MdSadiqMd This is a great solution! Approved. There were some conflicts with other merged PRs, so I am waiting on the CI rerun after those resolutions to complete the merge

@jlarson4
jlarson4 merged commit 984c165 into TransformerLensOrg:dev Sep 28, 2026
27 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[Proposal] One rejected fit discards an entire sparse-probing sweep, and controls carry no convergence diagnostics

2 participants