Files
OBLITERATUS/tests/test_whitened_svd_oracles.py
T

376 lines
14 KiB
Python

"""Reference-oracle tests for whitened SVD extraction.
## Test Context
- Code to test: `obliteratus/analysis/whitened_svd.py`
- Testing framework: pytest with torch CPU tensors
- Coverage target: repository default, minimum 80% Gate 3 target
- Test types needed: deterministic unit and metamorphic tests
- External dependencies to mock: none; these tests intentionally exercise the
real CPU tensor math and production `WhitenedSVDExtractor`
- Edge cases identified: paired sample permutation, feature-coordinate
permutation, common translation, SVD sign ambiguity, dtype normalization,
singular harmless covariance, zero signal, non-finite input, deterministic
replay, over-requested directions, and `extract_all_layers` intersection order
The static fixtures below are tiny deterministic activation matrices. Dynamic
fixture factories clone rows into the public list-of-tensors input shape so
tests cannot pass by mutating shared tensors across cases.
"""
from __future__ import annotations
from dataclasses import dataclass
import pytest
import torch
from obliteratus.analysis.whitened_svd import WhitenedSVDExtractor
REGULARIZATION_EPS = 1e-4
@dataclass(frozen=True)
class ReferenceWhitenedSVD:
directions: torch.Tensor
whitened_directions: torch.Tensor
singular_values: torch.Tensor
variance_explained: float
def _activation_pair_fixture(
dtype: torch.dtype = torch.float32,
) -> tuple[torch.Tensor, torch.Tensor]:
harmless = torch.tensor(
[
[-2.0, -1.0, 0.5, 1.0],
[-1.0, 0.0, -0.5, -1.0],
[0.0, 1.0, 1.5, 0.0],
[1.0, -2.0, -1.5, 2.0],
[2.0, 2.0, 0.0, -2.0],
],
dtype=dtype,
)
paired_delta = torch.tensor(
[
[1.25, -0.50, 0.75, 0.10],
[0.80, 0.25, -0.10, 0.40],
[1.60, -0.75, 0.35, -0.20],
[0.45, 0.90, -0.60, 0.55],
[1.10, -0.20, 0.95, -0.35],
],
dtype=dtype,
)
return harmless + paired_delta, harmless
def _singular_covariance_fixture() -> tuple[torch.Tensor, torch.Tensor]:
harmless = torch.tensor(
[
[-2.0, 0.0, 0.0, 1.0],
[-1.0, 0.0, 0.0, 1.0],
[0.0, 0.0, 0.0, 1.0],
[1.0, 0.0, 0.0, 1.0],
[2.0, 0.0, 0.0, 1.0],
],
)
harmful = harmless + torch.tensor([1.5, 0.0, 0.0, 0.0])
return harmful, harmless
def _as_public_samples(matrix: torch.Tensor) -> list[torch.Tensor]:
return [row.clone() for row in matrix]
def _reference_whitened_svd(
harmful: torch.Tensor,
harmless: torch.Tensor,
*,
n_directions: int,
min_variance_ratio: float = 0.0,
) -> ReferenceWhitenedSVD:
"""Independent tiny-tensor reference for the published mathematical contract."""
harmful = harmful.to(torch.float32).to(torch.float64)
harmless = harmless.to(torch.float32).to(torch.float64)
baseline_mean = harmless.mean(dim=0, keepdim=True)
centered_baseline = harmless - baseline_mean
covariance = centered_baseline.T.mm(centered_baseline) / max(harmless.shape[0] - 1, 1)
eigenvalues, eigenvectors = torch.linalg.eigh(covariance)
eigenvalues = torch.clamp(eigenvalues, min=0.0)
threshold = eigenvalues.max() * min_variance_ratio
kept = eigenvalues >= threshold
kept_values = eigenvalues[kept]
kept_vectors = eigenvectors[:, kept]
whitening = kept_vectors.mm(torch.diag(torch.rsqrt(kept_values + REGULARIZATION_EPS)))
whitened_delta = (harmful - baseline_mean).mm(whitening) - centered_baseline.mm(whitening)
_, singular_values, right_vectors_t = torch.linalg.svd(whitened_delta, full_matrices=False)
count = min(n_directions, whitened_delta.shape[0], whitened_delta.shape[1])
whitened_directions = right_vectors_t[:count]
inverse_whitening = kept_vectors.mm(torch.diag(torch.sqrt(kept_values + REGULARIZATION_EPS)))
original_directions = whitened_directions.mm(inverse_whitening.T)
original_directions = torch.nn.functional.normalize(original_directions, dim=1)
whitened_directions = torch.nn.functional.normalize(whitened_directions, dim=1)
selected_singular_values = singular_values[:count]
variance_explained = (
selected_singular_values.square().sum() / singular_values.square().sum().clamp(min=1e-12)
).item()
return ReferenceWhitenedSVD(
directions=original_directions,
whitened_directions=whitened_directions,
singular_values=selected_singular_values,
variance_explained=variance_explained,
)
def _row_space_projector(rows: torch.Tensor, *, tolerance: float = 1e-9) -> torch.Tensor:
rows = rows.to(torch.float64)
_, singular_values, right_vectors_t = torch.linalg.svd(rows, full_matrices=False)
basis = right_vectors_t[singular_values > tolerance]
return basis.T.mm(basis)
def _assert_same_subspace(left: torch.Tensor, right: torch.Tensor, *, atol: float = 3e-5) -> None:
assert torch.allclose(
_row_space_projector(left),
_row_space_projector(right),
atol=atol,
rtol=0,
)
def test_matches_independent_reference_not_raw_svd_or_identity_whitening() -> None:
harmful, harmless = _activation_pair_fixture()
result = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
).extract(_as_public_samples(harmful), _as_public_samples(harmless), n_directions=3)
reference = _reference_whitened_svd(harmful, harmless, n_directions=3)
_assert_same_subspace(result.directions, reference.directions)
_assert_same_subspace(result.whitened_directions, reference.whitened_directions)
assert result.singular_values.double() == pytest.approx(
reference.singular_values,
rel=3e-5,
abs=3e-5,
)
assert result.variance_explained == pytest.approx(reference.variance_explained, abs=3e-6)
raw_delta = harmful - harmless
_, _, raw_right_vectors_t = torch.linalg.svd(raw_delta, full_matrices=False)
raw_primary_alignment = torch.dot(result.directions[0], raw_right_vectors_t[0]).abs()
assert raw_primary_alignment < 0.95
def test_joint_sample_permutation_preserves_sign_invariant_refusal_subspace() -> None:
harmful, harmless = _activation_pair_fixture()
permutation = torch.tensor([3, 0, 4, 1, 2])
extractor = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
)
original = extractor.extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
permuted = extractor.extract(
_as_public_samples(harmful[permutation]),
_as_public_samples(harmless[permutation]),
3,
)
_assert_same_subspace(original.directions, permuted.directions)
assert permuted.singular_values == pytest.approx(original.singular_values, rel=3e-5, abs=3e-5)
assert permuted.variance_explained == pytest.approx(original.variance_explained, abs=3e-6)
def test_feature_coordinate_permutation_round_trips_through_inverse_mapping() -> None:
harmful, harmless = _activation_pair_fixture()
feature_permutation = torch.tensor([2, 0, 3, 1])
extractor = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
# The fixture has one mathematically null covariance dimension. Apply
# the production-default truncation policy so this metamorphic oracle
# compares the well-defined retained subspace across LAPACK backends.
min_variance_ratio=0.01,
)
original = extractor.extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
permuted = extractor.extract(
_as_public_samples(harmful[:, feature_permutation]),
_as_public_samples(harmless[:, feature_permutation]),
3,
)
mapped_back = torch.empty_like(permuted.directions)
mapped_back[:, feature_permutation] = permuted.directions
_assert_same_subspace(original.directions, mapped_back)
assert permuted.singular_values == pytest.approx(original.singular_values, rel=5e-5, abs=5e-5)
def test_common_translation_cannot_change_whitened_svd_oracle_values() -> None:
harmful, harmless = _activation_pair_fixture()
offset = torch.tensor([8.0, -3.0, 0.25, 11.0])
extractor = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
)
original = extractor.extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
translated = extractor.extract(
_as_public_samples(harmful + offset),
_as_public_samples(harmless + offset),
3,
)
_assert_same_subspace(original.directions, translated.directions)
_assert_same_subspace(original.whitened_directions, translated.whitened_directions)
assert translated.singular_values == pytest.approx(original.singular_values, rel=3e-5, abs=3e-5)
assert translated.variance_explained == pytest.approx(original.variance_explained, abs=3e-6)
@pytest.mark.parametrize("dtype", [torch.float32, torch.float64])
def test_float_inputs_follow_float32_output_policy_with_dtype_tolerances(
dtype: torch.dtype,
) -> None:
harmful, harmless = _activation_pair_fixture(dtype)
result = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
).extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
reference = _reference_whitened_svd(harmful, harmless, n_directions=3)
assert result.directions.dtype is torch.float32
assert result.whitened_directions.dtype is torch.float32
assert result.singular_values.dtype is torch.float32
_assert_same_subspace(result.directions, reference.directions, atol=4e-5)
assert result.singular_values.double() == pytest.approx(
reference.singular_values,
rel=4e-5,
abs=4e-5,
)
def test_singular_covariance_limits_over_requested_directions_to_effective_rank() -> None:
harmful, harmless = _singular_covariance_fixture()
result = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.01,
).extract(_as_public_samples(harmful), _as_public_samples(harmless), n_directions=5)
assert result.directions.shape == (1, 4)
assert result.whitened_directions.shape == (1, 1)
assert result.singular_values.shape == (1,)
assert torch.count_nonzero(result.directions[0, 1:].abs() > 1e-6) == 0
assert result.directions.norm() == pytest.approx(1.0)
assert result.variance_explained == pytest.approx(1.0)
assert result.effective_rank == pytest.approx(1.0, abs=1e-6)
def test_identical_harmful_and_harmless_inputs_are_rejected_as_no_refusal_signal() -> None:
_, harmless = _activation_pair_fixture()
extractor = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
)
with pytest.raises(ValueError, match="without activation difference"):
extractor.extract(_as_public_samples(harmless), _as_public_samples(harmless), 2)
@pytest.mark.parametrize(
("poisoned_side", "poisoned_value", "message"),
[
("harmful", float("nan"), "finite"),
("harmless", float("inf"), "finite"),
],
)
def test_non_finite_activation_values_are_rejected_before_linear_algebra(
poisoned_side: str,
poisoned_value: float,
message: str,
) -> None:
harmful, harmless = _activation_pair_fixture()
target = harmful if poisoned_side == "harmful" else harmless
target[1, 2] = poisoned_value
with pytest.raises(ValueError, match=message):
WhitenedSVDExtractor().extract(_as_public_samples(harmful), _as_public_samples(harmless), 1)
def test_deterministic_replay_returns_byte_stable_cpu_outputs() -> None:
harmful, harmless = _activation_pair_fixture()
extractor = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
)
first = extractor.extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
second = extractor.extract(_as_public_samples(harmful), _as_public_samples(harmless), 3)
assert torch.equal(first.directions, second.directions)
assert torch.equal(first.whitened_directions, second.whitened_directions)
assert torch.equal(first.singular_values, second.singular_values)
assert first.variance_explained == second.variance_explained
assert first.condition_number == second.condition_number
assert first.effective_rank == second.effective_rank
def test_extract_all_layers_returns_sorted_harmful_harmless_intersection_only() -> None:
base_harmful, base_harmless = _activation_pair_fixture()
harmful_by_layer = {
8: _as_public_samples(base_harmful + 0.25),
2: _as_public_samples(base_harmful),
5: _as_public_samples(base_harmful * 1.5),
}
harmless_by_layer = {
9: _as_public_samples(base_harmless),
5: _as_public_samples(base_harmless * 1.5),
2: _as_public_samples(base_harmless),
}
results = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
).extract_all_layers(harmful_by_layer, harmless_by_layer, n_directions=2)
assert list(results) == [2, 5]
assert [result.layer_idx for result in results.values()] == [2, 5]
assert all(result.directions.shape == (2, 4) for result in results.values())
def test_extract_all_layers_skips_missing_layers_without_stopping_later_matches() -> None:
base_harmful, base_harmless = _activation_pair_fixture()
harmful_by_layer = {
2: _as_public_samples(base_harmful),
5: _as_public_samples(base_harmful * 1.5),
8: _as_public_samples(base_harmful + 0.25),
}
harmless_by_layer = {
5: _as_public_samples(base_harmless * 1.5),
8: _as_public_samples(base_harmless),
}
results = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
).extract_all_layers(harmful_by_layer, harmless_by_layer, n_directions=2)
assert list(results) == [5, 8]
assert [result.layer_idx for result in results.values()] == [5, 8]
def test_extract_all_layers_uses_documented_default_direction_count() -> None:
base_harmful, base_harmless = _activation_pair_fixture()
harmful_by_layer = {5: _as_public_samples(base_harmful)}
harmless_by_layer = {5: _as_public_samples(base_harmless)}
results = WhitenedSVDExtractor(
regularization_eps=REGULARIZATION_EPS,
min_variance_ratio=0.0,
).extract_all_layers(harmful_by_layer, harmless_by_layer)
assert list(results) == [5]
assert results[5].directions.shape == (4, 4)