Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
16 changes: 11 additions & 5 deletions docs/source/content/sparse_probing.md
Original file line number Diff line number Diff line change
Expand Up @@ -59,8 +59,11 @@ $$

The selected support contains the $k$ largest $|s_j|$. Equal scores are resolved by increasing
feature index. `preprocess="none"` fits the selected raw coordinates. With
`preprocess="standardize"`, selected columns are centered and scaled using training statistics;
zero-variance columns receive scale one. The same transform is then applied to held-out values.
`preprocess="standardize"`, selected columns are centered and scaled using training statistics.
Each column's scale is its training standard deviation raised to at least `std_floor` (default
`1e-3`), so a near-constant column is not amplified to unit scale; zero-variance columns receive
scale one, and `std_floor=0` disables the floor. The same transform is then applied to held-out
values.

Because L2 regularization is scale-sensitive, preprocessing can change the fitted probe and
the resulting k-curve. A sweep therefore fixes preprocessing and L2 strength across every k.
Expand All @@ -84,8 +87,9 @@ Accuracy, precision, recall, F1, and all four confusion counts are returned; pre
zero when its denominator is zero. F1 is the primary sparse-probing metric.

Feature-score reductions use float64 for float64 inputs and float32 otherwise. Selected matrices
move to CPU float64, where LBFGS (at most `max_iter` iterations, with a fixed internal gradient
stop) is followed by up to `max_refinement_steps` damped Newton steps on the `(k+1)`-square
move to CPU float64, where LBFGS (at most `max_iter` iterations and `max_iter * 5 // 4` function
Comment thread
jlarson4 marked this conversation as resolved.
Outdated
evaluations, with a fixed internal gradient stop) is followed by up to `max_refinement_steps`
damped Newton steps on the `(k+1)`-square
Hessian of the objective. A fit is accepted only when the Newton decrement $\tfrac{1}{2}
g^\top H^{-1} g$, an estimate of the objective gap to the optimum in nats, is at most
`decrement_tolerance` (default `1e-12`, which must lie in `(0, 1)`); the decrement is checked
Expand Down Expand Up @@ -163,4 +167,6 @@ with [reference code](https://github.com/wesg52/sparse-probing-paper).

TransformerLens intentionally adds stratification, stable tie-breaking, explicit objective and
convergence diagnostics, and deterministic controls. Its optional centered standardization and
Torch LBFGS solver are not exact reproductions of the reference implementation.
Torch LBFGS solver are not exact reproductions of the reference implementation. The default
`std_floor=1e-3` matches the reference's standard-deviation floor, except that zero-variance
columns keep scale one.
48 changes: 47 additions & 1 deletion tests/unit/tools/test_sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,43 @@ def test_standardization_is_train_only_and_heldout_values_do_not_change_selectio
assert first.objective == second.objective


def test_standardization_floor_near_constant_column_scale():
generator = torch.Generator().manual_seed(0)
features = torch.randn(200, 8, generator=generator)
features[:, 0] = 1.0 + torch.randn(200, generator=generator) * 1e-6
labels = (torch.rand(200, generator=generator) < 0.5).long()

# The near-constant column has the smallest |score|, so only k=8 puts it in the support.
result = fit_sparse_probe(features, labels, k=8, preprocess="standardize", seed=0)
near_constant = result.selected_features.tolist().index(0)

assert not result.constant_features[near_constant]
assert result.preprocess_scale[near_constant].item() == 1e-3
assert result.std_floor == 1e-3


def test_std_floor_keeps_constant_columns_at_scale_one_and_zero_disables_it():
generator = torch.Generator().manual_seed(0)
features = torch.randn(200, 8, generator=generator)
features[:, 0] = 1.0 + torch.randn(200, generator=generator) * 1e-6
features[:, 1] = 5.0
labels = (torch.rand(200, generator=generator) < 0.5).long()

floored = fit_sparse_probe(features, labels, k=8, preprocess="standardize", seed=0)
unfloored = fit_sparse_probe(
features, labels, k=8, preprocess="standardize", std_floor=0, seed=0
)

support = floored.selected_features.tolist()
near_constant, constant = support.index(0), support.index(1)
selected_train = features[floored.train_indices][:, floored.selected_features].double()
assert floored.constant_features[constant]
assert floored.preprocess_scale[constant].item() == 1.0
assert unfloored.preprocess_scale[near_constant].item() == pytest.approx(
selected_train[:, near_constant].std(correction=0).item()
)


def test_none_preprocessing_has_identity_metadata_and_constant_tie_order():
features = torch.zeros(20, 5)
labels = torch.arange(20) % 2
Expand Down Expand Up @@ -529,6 +566,9 @@ def test_binary_metrics_zero_division_policy():
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"class_weight": "bad"}, "class_weight"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"seed": -1}, "seed"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"preprocess": "bad"}, "preprocess"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"std_floor": -1}, "std_floor"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"std_floor": float("nan")}, "std_floor"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"std_floor": True}, "std_floor"),
(torch.ones(4, 2), torch.tensor([0, 1, 0, 1]), {"max_iter": 0}, "max_iter"),
(
torch.ones(4, 2),
Expand Down Expand Up @@ -716,6 +756,7 @@ def test_label_shuffle_control_is_scored_against_the_true_heldout_labels():
test_fraction=0.3,
positive_label=1,
preprocess="none",
std_floor=1e-3,
class_weight="balanced",
l2_strength=1e-2,
seed=seed,
Expand All @@ -737,7 +778,12 @@ def test_label_shuffle_control_is_scored_against_the_true_heldout_labels():
assert torch.equal(support, control.supports[0])

train_features, test_features, *_ = _selected_data(
validated.features, support, train_indices, test_indices, validated.preprocess
validated.features,
support,
train_indices,
test_indices,
validated.preprocess,
validated.std_floor,
)
fit = _fit_logistic(
train_features,
Expand Down
33 changes: 30 additions & 3 deletions transformer_lens/tools/analysis/sparse_probing.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@ class SparseProbeResult:
test_positive_count: int
test_negative_count: int
preprocess: PreprocessMode
std_floor: float
class_weight: ClassWeightMode
l2_strength: float
test_fraction: float
Expand Down Expand Up @@ -113,6 +114,7 @@ class _ValidatedInputs:
k: int
test_fraction: float
preprocess: PreprocessMode
std_floor: float
class_weight: ClassWeightMode
l2_strength: float
seed: int
Expand Down Expand Up @@ -157,6 +159,7 @@ def _validate_inputs(
test_fraction: int | float,
positive_label: int | bool,
preprocess: str,
std_floor: int | float,
class_weight: str | None,
l2_strength: int | float,
seed: int,
Expand Down Expand Up @@ -217,6 +220,12 @@ def _validate_inputs(
raise ValueError(f"test_fraction must be a finite real in (0, 1), got {test_fraction!r}")
if preprocess not in ("none", "standardize"):
raise ValueError(f"preprocess must be 'none' or 'standardize', got {preprocess!r}")
if isinstance(std_floor, bool) or not isinstance(std_floor, (int, float)):
raise ValueError(f"std_floor must be a finite nonnegative real, got {std_floor!r}")
validated_std_floor = float(std_floor)
if not math.isfinite(validated_std_floor) or validated_std_floor < 0:
raise ValueError(f"std_floor must be a finite nonnegative real, got {std_floor!r}")

validated_preprocess = cast(PreprocessMode, preprocess)
if class_weight not in ("balanced", None):
raise ValueError(f"class_weight must be 'balanced' or None, got {class_weight!r}")
Expand All @@ -239,6 +248,7 @@ def _validate_inputs(
k=validated_k,
test_fraction=validated_fraction,
preprocess=validated_preprocess,
std_floor=validated_std_floor,
class_weight=validated_class_weight,
l2_strength=validated_l2,
seed=seed,
Expand Down Expand Up @@ -284,6 +294,7 @@ def _selected_data(
train_indices: torch.Tensor,
test_indices: torch.Tensor,
preprocess: PreprocessMode,
std_floor: float,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
selected_device = selected_features.to(device=features.device)
train_device = train_indices.to(device=features.device)
Expand All @@ -304,7 +315,8 @@ def _selected_data(
constant = raw_scale == 0
if preprocess == "standardize":
mean = train.mean(dim=0)
scale = torch.where(constant, torch.ones_like(raw_scale), raw_scale)
# Zero-variance columns keep scale one: flooring them would inflate held-out deviations.
Comment thread
jlarson4 marked this conversation as resolved.
Outdated
scale = torch.where(constant, torch.ones_like(raw_scale), raw_scale.clamp(min=std_floor))
return (train - mean) / scale, (test - mean) / scale, mean, scale, constant
mean = torch.zeros(train.shape[1], dtype=torch.float64)
scale = torch.ones(train.shape[1], dtype=torch.float64)
Expand Down Expand Up @@ -563,6 +575,7 @@ def _fit_result(
train_indices,
test_indices,
validated.preprocess,
validated.std_floor,
)
train_labels = validated.canonical_labels[train_indices]
fit = _fit_logistic(
Expand Down Expand Up @@ -594,6 +607,7 @@ def _fit_result(
test_positive_count=int(test_labels.sum().item()),
test_negative_count=int((~test_labels).sum().item()),
preprocess=validated.preprocess,
std_floor=validated.std_floor,
class_weight=validated.class_weight,
l2_strength=validated.l2_strength,
test_fraction=validated.test_fraction,
Expand Down Expand Up @@ -625,6 +639,7 @@ def _fit_control(
train_indices,
test_indices,
validated.preprocess,
validated.std_floor,
)
fit = _fit_logistic(
train_features,
Expand Down Expand Up @@ -674,6 +689,7 @@ def fit_sparse_probe(
test_fraction: int | float = 0.3,
positive_label: int | bool = 1,
preprocess: str = "none",
std_floor: int | float = 1e-3,
class_weight: str | None = "balanced",
l2_strength: int | float = 1e-2,
seed: int = 0,
Expand All @@ -694,10 +710,14 @@ def fit_sparse_probe(
test_fraction: Requested held-out fraction within each class.
positive_label: Label defining the positive class and score sign.
preprocess: ``"none"`` or train-only ``"standardize"``.
std_floor: Lower bound on a non-constant column's ``"standardize"`` scale, so a
near-constant column is not amplified to unit scale; ``0`` disables it.
class_weight: ``"balanced"`` or ``None`` for unweighted BCE.
l2_strength: Positive coefficient penalty in the logistic objective.
seed: Local CPU-generator seed used only for the stratified split.
max_iter: Maximum LBFGS iterations before Newton refinement.
max_iter: Maximum LBFGS iterations before Newton refinement. LBFGS also stops after
``max_iter * 5 // 4`` function evaluations (``stop_reason="max_eval"``), which
equals ``max_iter`` when ``max_iter <= 3``.
max_refinement_steps: Maximum damped Newton steps after LBFGS; zero only checks.
decrement_tolerance: Largest accepted Newton decrement ``g^T H^-1 g / 2``, an
estimate of the objective gap to the optimum in nats; a larger gap raises.
Expand All @@ -717,6 +737,7 @@ def fit_sparse_probe(
test_fraction=test_fraction,
positive_label=positive_label,
preprocess=preprocess,
std_floor=std_floor,
class_weight=class_weight,
l2_strength=l2_strength,
seed=seed,
Expand Down Expand Up @@ -752,6 +773,7 @@ def sweep_sparse_probe(
test_fraction: int | float = 0.3,
positive_label: int | bool = 1,
preprocess: str = "none",
std_floor: int | float = 1e-3,
Comment thread
jlarson4 marked this conversation as resolved.
class_weight: str | None = "balanced",
l2_strength: int | float = 1e-2,
n_random_subsets: int = 0,
Expand All @@ -777,12 +799,16 @@ def sweep_sparse_probe(
test_fraction: Requested held-out fraction within each class.
positive_label: Label defining the positive class and score sign.
preprocess: ``"none"`` or train-only ``"standardize"``.
std_floor: Lower bound on a non-constant column's ``"standardize"`` scale, shared by
every fit; ``0`` disables it.
class_weight: ``"balanced"`` or ``None``, shared by every fit.
l2_strength: Positive coefficient penalty shared by every fit.
n_random_subsets: Random-coordinate control fits per sparsity level.
n_label_shuffles: Shuffled-training-label control fits per sparsity level.
seed: Local CPU-generator seed for splitting and controls.
max_iter: Maximum LBFGS iterations per fit before Newton refinement.
max_iter: Maximum LBFGS iterations per fit before Newton refinement. LBFGS also stops
after ``max_iter * 5 // 4`` function evaluations (``stop_reason="max_eval"``),
which equals ``max_iter`` when ``max_iter <= 3``.
max_refinement_steps: Maximum damped Newton steps per fit; zero only checks.
decrement_tolerance: Largest accepted Newton decrement ``g^T H^-1 g / 2`` per fit,
an estimate of the objective gap to the optimum in nats; a larger gap raises.
Expand Down Expand Up @@ -812,6 +838,7 @@ def sweep_sparse_probe(
test_fraction=test_fraction,
positive_label=positive_label,
preprocess=preprocess,
std_floor=std_floor,
class_weight=class_weight,
l2_strength=l2_strength,
seed=seed,
Expand Down
Loading