Files
OBLITERATUS/tests/test_leace.py
T

100 lines
3.3 KiB
Python

"""Numerical and error-path tests for Fisher/LEACE direction extraction."""
from __future__ import annotations
import pytest
import torch
from obliteratus.analysis.leace import LEACEExtractor
def _activations():
harmful = [
torch.tensor([2.0, 0.0, 1.0]),
torch.tensor([3.0, 1.0, 0.0]),
torch.tensor([4.0, -1.0, 1.0]),
]
harmless = [
torch.tensor([-2.0, 0.0, 1.0]),
torch.tensor([-3.0, 1.0, 0.0]),
torch.tensor([-4.0, -1.0, 1.0]),
]
return harmful, harmless
def test_extract_returns_normalized_reproducible_diagnostics():
harmful, harmless = _activations()
result = LEACEExtractor().extract(harmful, harmless, layer_idx=7)
assert result.layer_idx == 7
assert result.direction.norm().item() == pytest.approx(1.0)
assert result.direction[0].abs().item() > 0.99
assert result.generalized_eigenvalue > 0
assert result.mean_diff_norm == pytest.approx(6.0)
assert result.erasure_loss > 0
assert result.within_class_condition >= 1
def test_extract_accepts_singleton_sequence_axis_and_shrinkage():
harmful, harmless = _activations()
harmful_3d = [value.unsqueeze(0) for value in harmful]
harmless_3d = [value.unsqueeze(0) for value in harmless]
result = LEACEExtractor(shrinkage=0.5).extract(harmful_3d, harmless_3d)
assert result.direction.shape == (3,)
assert torch.isfinite(result.direction).all()
def test_degenerate_classes_return_zero_direction():
samples = [torch.ones(3), torch.ones(3)]
result = LEACEExtractor().extract(samples, samples)
assert torch.equal(result.direction, torch.zeros(3))
assert result.generalized_eigenvalue == 0
def test_solver_failure_uses_least_squares(monkeypatch):
harmful, harmless = _activations()
def fail_solve(*_args, **_kwargs):
raise torch.linalg.LinAlgError("fixture")
monkeypatch.setattr(torch.linalg, "solve", fail_solve)
result = LEACEExtractor().extract(harmful, harmless)
assert result.direction.norm().item() == pytest.approx(1.0)
def test_condition_failure_is_reported_as_infinite(monkeypatch):
harmful, harmless = _activations()
original_solve = torch.linalg.solve
monkeypatch.setattr(
torch.linalg,
"eigvalsh",
lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError("fixture")),
)
monkeypatch.setattr(torch.linalg, "solve", original_solve)
result = LEACEExtractor().extract(harmful, harmless)
assert result.within_class_condition == float("inf")
def test_extract_all_layers_skips_unpaired_layers_and_sorts():
harmful, harmless = _activations()
results = LEACEExtractor().extract_all_layers(
{2: harmful, 1: harmful, 3: harmful},
{1: harmless, 2: harmless},
)
assert list(results) == [1, 2]
assert results[1].layer_idx == 1
def test_compare_with_diff_of_means_handles_regular_and_degenerate_difference():
harmful, harmless = _activations()
result = LEACEExtractor().extract(harmful, harmless)
comparison = LEACEExtractor.compare_with_diff_of_means(
result,
torch.stack(harmful).mean(0),
torch.stack(harmless).mean(0),
)
assert comparison["cosine_similarity"] == pytest.approx(1.0)
degenerate = LEACEExtractor.compare_with_diff_of_means(
result, torch.zeros(3), torch.zeros(3),
)
assert degenerate["cosine_similarity"] == 0