Conversation
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
left a comment
There was a problem hiding this comment.
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?
|
@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:
The public |
|
Looks great @akshpatel0910, thank you for putting this together. Approved and merged |
Description
Adds tie-aware
roc_aucandaverage_precisiontoSparseProbeMetricsand propagates both toSparseProbeControlas 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 exactly0.0, and the all-positive prediction scores F12p/(1+p)(0.700 on the issue's fixture) despite carrying no information. The new metrics are computed from the held-out logits_binary_metricsalready 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 ratep.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_metricscalled directly can hit it.roc_aucisNaNwhen either class is absent (there is no positive-negative pair to rank).average_precisionisNaNwith no positives and1.0with no negatives (every threshold has precision one, the usual AP definition). This is stated in theSparseProbeMetricsdocstring.One note on parity with the reference:
experiments/metrics.py::get_binary_cls_perf_metricscomputestest_average_precisionfromy_pred(hard predictions) rather thany_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_aucin the sweep example.Fixes #1814
Type of change
SparseProbeMetricsandSparseProbeControlgain 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; assertsconstant_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_undefinedandtest_binary_metrics_all_positive_labels_give_nan_roc_auc_and_unit_average_precision: single-class policy, one test per case.Checklist:
Local checks (Python 3.12, CPU):
pytest tests/unit -m "not slow": 6702 passed, 56 skipped, 6 xfailedpytest tests/unit/tools/test_sparse_probing.py: 80 passedpytest transformer_lens/(docstring tests): 13 passed, 24 skippeduv run mypy .: no issuespycln --check,isort --check-only,black --check: clean