Skip to content

Add tie-aware roc_auc and average_precision to sparse probe metrics - #1820

Merged
jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
akshpatel0910:sparse-probe-roc-auc-average-precision
Sep 28, 2026
Merged

jlarson4 merged 2 commits into
TransformerLensOrg:devfrom
akshpatel0910:sparse-probe-roc-auc-average-precision

Conversation

@akshpatel0910

@akshpatel0910 akshpatel0910 commented Sep 25, 2026 •

Copy link
Copy Markdown

Description

Adds tie-aware roc_auc and average_precision to SparseProbeMetrics and propagates both to SparseProbeControl as per-repeat distributions.

Every existing metric is read at the single threshold logit >= 0. When the 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 the all-positive prediction scores F1 2p/(1+p) (0.700 on the issue's fixture) despite carrying no information. The new metrics are computed from the held-out logits _binary_metrics already has, with no threshold:

  • roc_auc: Mann-Whitney U with average ranks for tied logits, so the degenerate probe scores 0.5 (not 1.0).
  • average_precision: step-wise AP over distinct logits; tied logits form a single threshold, so the degenerate probe scores the positive rate p.

Single-class policy: the stratified split always holds out at least one example of each class, so this never arises from fit_sparse_probe / sweep_sparse_probe; _binary_metrics called directly can hit it. roc_auc is NaN when either class is absent (there is no positive-negative pair to rank). average_precision is NaN with no positives and 1.0 with no negatives (every threshold has precision one, the usual AP definition). This is stated in the SparseProbeMetrics docstring.

One note on parity with the reference: experiments/metrics.py::get_binary_cls_perf_metrics computes test_average_precision from y_pred (hard predictions) rather than y_score. This PR scores average precision from the logits, which is the threshold-free quantity the issue asks for.

No new dependency: tests check against a brute-force pairwise / per-threshold reference instead of scikit-learn.

The sparse probing guide now says that F1 is one operating point, names the dead-coordinate case, and shows roc_auc in the sweep example.

Fixes #1814

Type of change

  • New feature (non-breaking change which adds functionality)
  • This change requires a documentation update

SparseProbeMetrics and SparseProbeControl gain two fields each. Code that reads them is unaffected; code that constructs them directly (only the module itself does) would need the new fields.

Tests

  • test_dead_coordinate_probe_scores_chance_on_threshold_free_metrics: the fixture from the issue; asserts constant_features == [True], zero coefficient and intercept, F1 == 2p/(1+p) (≈ 0.700), roc_auc == 0.5, average_precision == p.
  • test_threshold_free_metrics_match_a_brute_force_tie_aware_reference: 5 seeds of heavily tied logits against pairwise-win ROC-AUC and a per-threshold AP loop.
  • test_threshold_free_metrics_match_hand_computed_values: perfect, inverted, all-tied, and cross-class-tie cases.
  • test_binary_metrics_all_negative_labels_leave_both_threshold_free_scores_undefined and test_binary_metrics_all_positive_labels_give_nan_roc_auc_and_unit_average_precision: single-class policy, one test per case.
  • Existing control tests extended: new fields propagate, are deterministic, are empty with zero repeats, match a re-scored label-shuffle fit, and the strong planted feature beats both controls on ROC-AUC.

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

Local checks (Python 3.12, CPU):

  • pytest tests/unit -m "not slow": 6702 passed, 56 skipped, 6 xfailed
  • pytest tests/unit/tools/test_sparse_probing.py: 80 passed
  • pytest transformer_lens/ (docstring tests): 13 passed, 24 skipped
  • uv run mypy .: no issues
  • pycln --check, isort --check-only, black --check: clean

F1 and the other threshold metrics are read at logit >= 0 only, so a probe
on a dead (constant) coordinate that predicts every held-out example positive
scores F1 = 2p/(1+p). Add threshold-free, tie-aware ROC-AUC (Mann-Whitney
with average ranks) and average precision (tied logits form one threshold)
to SparseProbeMetrics, propagate both to SparseProbeControl, and document
the degenerate case in the sparse probing guide.

Fixes TransformerLensOrg#1814

@koriyoshi2041 koriyoshi2041 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The tie-aware implementations match the hand-computed cases, and the focused sparse-probing suite passes locally (80/80). One contract detail should be made explicit before merge: the PR description says both metrics return NaN when either class is absent, and test_binary_metrics_threshold_free_scores_are_nan_without_both_classes has the same name, but _average_precision returns 1.0 for an all-positive input (the test intentionally expects that). That behavior is defensible and matches the usual AP definition; could you update the description/test name (or change the implementation if symmetric NaN was intended) so the documented single-class policy matches the API?

@akshpatel0910

Copy link
Copy Markdown
Author

@koriyoshi2041 Thanks for catching that! I kept the implementation, since AP = 1.0 without negatives follows the usual definition, and made the contract explicit instead:

  • Split the test into test_binary_metrics_all_negative_labels_leave_both_threshold_free_scores_undefined and test_binary_metrics_all_positive_labels_give_nan_roc_auc_and_unit_average_precision, so each name matches what it asserts.
  • Documented the policy in the SparseProbeMetrics docstring: roc_auc is NaN if either class is absent; average_precision is NaN without positives and 1.0 without negatives.
  • Corrected the PR description to match.

The public fit_sparse_probe / sweep_sparse_probe paths can't hit this, since the stratified split always holds out both classes. Changes are in d6a09be.

@jlarson4

Copy link
Copy Markdown
Collaborator

Looks great @akshpatel0910, thank you for putting this together. Approved and merged

@jlarson4
jlarson4 merged commit 0d2b606 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] Sparse probing reports only threshold metrics, so a zero-information probe scores F1 0.700

3 participants