Files
OBLITERATUS/tests/test_projection_math_contracts.py
T

795 lines
31 KiB
Python

"""Reference-oracle contracts for pure projection and orthogonalization math."""
from __future__ import annotations
import math
from types import SimpleNamespace
import pytest
import torch
from obliteratus.abliterate import AbliterationPipeline
from obliteratus.analysis.numerical_contracts import (
orthogonalize_subspace_rows,
project_weight_against_direction,
remove_harmless_principal_components,
residualize_against_shield_atoms,
select_projection_coefficients,
)
def _canonical_rows(rows: torch.Tensor) -> torch.Tensor:
result = rows.clone()
for idx in range(result.shape[0]):
pivot = result[idx].abs().argmax()
if result[idx, pivot] < 0:
result[idx] = -result[idx]
return result
def _reference_gram_schmidt(rows: torch.Tensor) -> torch.Tensor:
basis: list[torch.Tensor] = []
for row in rows.to(dtype=torch.float64):
residual = row.clone()
for prev in basis:
residual = residual - (residual @ prev) * prev
norm = residual.norm()
if norm > 1e-8:
basis.append(residual / norm)
return torch.stack(basis)[: rows.shape[0]].to(dtype=rows.dtype, device=rows.device)
def _reference_projection(weight: torch.Tensor, direction: torch.Tensor, scale: float) -> torch.Tensor:
work = weight.to(dtype=torch.float64)
d = direction.reshape(-1, 1).to(dtype=torch.float64, device=weight.device)
d_norm = d.norm()
if d_norm < 1e-8:
return weight.clone()
d = d / d_norm
if weight.shape[-1] == d.shape[0]:
coeff = work @ d
return (work - d.T * (scale * coeff)).to(dtype=weight.dtype)
coeff = d.T @ work
return (work - (scale * d) * coeff).to(dtype=weight.dtype)
def test_orthogonalize_matches_reference_and_preserves_primary_orientation():
subspace = torch.tensor(
[[2.0, 0.0, 0.0], [1.0, 3.0, 0.0], [1.0, 1.0, 4.0]],
dtype=torch.float64,
)
actual = orthogonalize_subspace_rows(subspace)
expected = _reference_gram_schmidt(subspace)
assert actual.dtype == subspace.dtype
assert torch.allclose(actual @ actual.T, torch.eye(3, dtype=torch.float64), atol=1e-12)
assert torch.allclose(_canonical_rows(actual), _canonical_rows(expected), atol=1e-12)
assert actual[0] @ subspace[0] > 0
def test_orthogonalize_returns_degenerate_inputs_without_allocating_new_tensor():
single_row = torch.tensor([[1.0, 2.0, 3.0]], dtype=torch.float64)
empty = torch.empty((2, 0), dtype=torch.float64)
assert orthogonalize_subspace_rows(single_row) is single_row
assert orthogonalize_subspace_rows(empty) is empty
def test_orthogonalize_non_finite_subspace_is_explicit_noop():
subspace = torch.tensor([[1.0, 0.0], [float("nan"), 1.0]], dtype=torch.float64)
actual = orthogonalize_subspace_rows(subspace)
assert actual is subspace
assert torch.isnan(actual[1, 0])
def test_orthogonalization_preserves_float64_precision_for_nearly_collinear_rows():
subspace = torch.tensor(
[[1.0, 1e-6, 0.0], [1.0, 0.0, 1e-6], [0.0, 1.0, 1.0]],
dtype=torch.float64,
)
actual = orthogonalize_subspace_rows(subspace)
expected = _reference_gram_schmidt(subspace)
assert actual.dtype == torch.float64
assert torch.allclose(_canonical_rows(actual), _canonical_rows(expected), atol=1e-10)
def test_orthogonalize_two_half_precision_rows_uses_stable_compute_dtype():
subspace = torch.tensor(
[[1.0, 1.0, 0.0], [1.0, 0.0, 1.0]],
dtype=torch.float16,
)
actual = orthogonalize_subspace_rows(subspace)
expected = _reference_gram_schmidt(subspace)
assert actual is not subspace
assert actual.dtype == torch.float16
assert torch.allclose(_canonical_rows(actual), _canonical_rows(expected), atol=1e-3)
gram = actual.float() @ actual.float().T
assert torch.allclose(gram, torch.eye(2), atol=1e-3)
def test_integer_projection_uses_supported_float_compute_then_restores_weight_dtype():
weight = torch.tensor([[2, 0], [0, 2]], dtype=torch.int64)
direction = torch.tensor([1, 0], dtype=torch.int64)
projected = project_weight_against_direction(weight, direction)
assert projected.projected is True
assert projected.weight.dtype == torch.int64
assert torch.equal(projected.weight, torch.tensor([[0, 0], [0, 2]], dtype=torch.int64))
def test_projection_full_removal_is_idempotent_and_orthogonal_to_direction():
weight = torch.tensor([[3.0, 4.0, 0.0], [1.0, -2.0, 2.0]], dtype=torch.float64)
direction = torch.tensor([0.6, 0.8, 0.0], dtype=torch.float64)
first = project_weight_against_direction(weight, direction, regularization=0.0)
second = project_weight_against_direction(first.weight, direction, regularization=0.0)
assert first.projected
assert torch.allclose(first.weight @ direction, torch.zeros(2, dtype=torch.float64), atol=1e-12)
assert torch.allclose(first.weight, second.weight, atol=1e-12)
assert torch.allclose(first.weight, _reference_projection(weight, direction, scale=1.0), atol=1e-12)
def test_projection_normalizes_non_unit_directions_before_applying_formula():
weight = torch.tensor([[3.0, 4.0], [-5.0, 6.0]], dtype=torch.float64)
unit = project_weight_against_direction(weight, torch.tensor([1.0, 0.0], dtype=torch.float64))
non_unit = project_weight_against_direction(weight, torch.tensor([2.0, 0.0], dtype=torch.float64))
assert unit.projected
assert non_unit.projected
assert torch.allclose(non_unit.weight, unit.weight, atol=1e-12)
assert torch.allclose(non_unit.weight, _reference_projection(weight, torch.tensor([2.0, 0.0]), 1.0))
def test_projection_zero_direction_is_deterministic_no_op():
weight = torch.tensor([[3.0, 4.0], [-5.0, 6.0]], dtype=torch.float64)
direction = torch.zeros(2, dtype=torch.float64)
first = project_weight_against_direction(weight, direction)
second = project_weight_against_direction(weight, direction)
assert not first.projected
assert first.layout is None
assert torch.equal(first.weight, weight)
assert torch.equal(second.weight, first.weight)
assert first.coefficient_norm_sq == 0.0
def test_projection_projects_at_exact_tiny_direction_threshold():
weight = torch.tensor([[3.0, 4.0]], dtype=torch.float64)
direction = torch.tensor([1e-8, 0.0], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction)
assert projected.projected is True
assert projected.layout == "standard"
assert torch.allclose(projected.weight, torch.tensor([[0.0, 4.0]], dtype=torch.float64), atol=1e-12)
def test_projection_zero_direction_metadata_uses_strict_false_flag():
result = project_weight_against_direction(
torch.tensor([[3.0, 4.0]], dtype=torch.float64),
torch.zeros(2, dtype=torch.float64),
)
assert result.projected is False
assert result.weight is not None
assert result.coefficient_norm_sq == 0.0
assert result.layout is None
@pytest.mark.parametrize("projection_row_fraction", [0.0, -0.1, 1.01])
def test_select_projection_coefficients_rejects_invalid_fraction(projection_row_fraction):
coeff = torch.tensor([[1.0], [2.0]], dtype=torch.float64)
with pytest.raises(ValueError, match=r"projection_row_fraction must be in"):
select_projection_coefficients(coeff, projection_row_fraction)
def test_select_projection_coefficients_empty_and_singleton_inputs_are_noops():
empty = torch.empty((0, 1), dtype=torch.float64)
singleton = torch.tensor([[3.0]], dtype=torch.float64)
assert select_projection_coefficients(empty, 0.5) is empty
assert select_projection_coefficients(singleton, 0.5) is singleton
def test_projection_unsupported_layout_returns_complete_noop_metadata_clone():
weight = torch.tensor([[1.0, 2.0, 3.0]], dtype=torch.float64)
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction)
assert projected.projected is False
assert projected.layout is None
assert projected.coefficient_norm_sq == 0.0
assert projected.weight is not weight
assert torch.equal(projected.weight, weight)
def test_projection_standard_layout_non_finite_coefficients_fail_closed():
huge = torch.finfo(torch.float32).max
weight = torch.tensor([[huge, huge]], dtype=torch.float32)
direction = torch.tensor([1.0, 1.0], dtype=torch.float32)
projected = project_weight_against_direction(weight, direction)
assert projected.projected is False
assert projected.layout == "standard"
assert projected.coefficient_norm_sq == 0.0
assert torch.equal(projected.weight, weight)
def test_projection_transposed_layout_non_finite_coefficients_fail_closed():
huge = torch.finfo(torch.float32).max
weight = torch.tensor([[huge], [huge]], dtype=torch.float32)
direction = torch.tensor([1.0, 1.0], dtype=torch.float32)
projected = project_weight_against_direction(weight, direction)
assert projected.projected is False
assert projected.layout == "transposed"
assert projected.coefficient_norm_sq == 0.0
assert torch.equal(projected.weight, weight)
def test_projection_preserves_orthogonal_coordinates_and_contracts_norm_without_restore():
direction = torch.tensor([1.0, 0.0, 0.0])
orthogonal_probe = torch.tensor([0.0, 2.0, -1.0])
weight = torch.tensor([[5.0, 2.0, -1.0], [-3.0, 4.0, 7.0]])
projected = project_weight_against_direction(weight, direction, norm_preserve=False)
assert torch.allclose(projected.weight @ direction, torch.zeros(2))
assert torch.allclose(projected.weight @ orthogonal_probe, weight @ orthogonal_probe)
assert projected.weight.norm() <= weight.norm()
def test_norm_preservation_uses_cap_when_projection_would_amplify_too_much():
weight = torch.tensor([[100.0, 1.0], [0.0, 0.0]])
direction = torch.tensor([1.0, 0.0])
projected = project_weight_against_direction(
weight,
direction,
norm_preserve=True,
max_norm_ratio=1.10,
)
assert projected.projected
assert torch.isclose(projected.weight.norm(), torch.tensor(1.10), atol=1e-6)
def test_norm_preservation_keeps_zero_projection_without_restoration():
weight = torch.tensor([[3.0, 0.0]], dtype=torch.float64)
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction, norm_preserve=True)
assert projected.projected is True
assert projected.layout == "standard"
assert projected.coefficient_norm_sq == 9.0
assert torch.allclose(projected.weight, torch.zeros_like(weight), atol=0.0, rtol=0.0)
assert torch.isfinite(projected.weight).all()
def test_norm_preservation_at_max_ratio_boundary_preserves_original_norm():
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
weight = torch.tensor([[math.sqrt(21.0), 10.0]], dtype=torch.float64)
projected = project_weight_against_direction(
weight,
direction,
norm_preserve=True,
max_norm_ratio=1.10,
)
assert projected.projected
assert torch.allclose(projected.weight.norm(), weight.norm(), atol=1e-12)
def test_transposed_norm_preservation_reports_removed_coefficient_energy():
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
weight = torch.tensor([[3.0, 4.0, 0.0], [10.0, 20.0, 30.0]], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction, norm_preserve=True)
assert projected.projected is True
assert projected.layout == "transposed"
assert projected.coefficient_norm_sq == 25.0
assert torch.allclose(projected.weight[0], torch.zeros(3, dtype=torch.float64), atol=1e-12)
def test_transposed_projection_without_norm_preservation_reports_zero_metadata_energy():
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
weight = torch.tensor([[3.0, 4.0, 0.0], [10.0, 20.0, 30.0]], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction, norm_preserve=False)
assert projected.projected is True
assert projected.layout == "transposed"
assert projected.coefficient_norm_sq == 0.0
assert torch.allclose(projected.weight[0], torch.zeros(3, dtype=torch.float64), atol=1e-12)
def test_projection_supports_standard_and_transposed_layouts():
direction = torch.tensor([1.0, 0.0])
standard = torch.tensor([[3.0, 4.0], [5.0, 6.0], [7.0, 8.0]])
transposed = standard.T.contiguous()
standard_result = project_weight_against_direction(standard, direction)
transposed_result = project_weight_against_direction(transposed, direction)
assert torch.allclose(standard_result.weight[:, 0], torch.zeros(3))
assert torch.allclose(transposed_result.weight[0, :], torch.zeros(3))
assert torch.allclose(standard_result.weight, _reference_projection(standard, direction, 1.0))
assert torch.allclose(transposed_result.weight, _reference_projection(transposed, direction, 1.0))
def test_projection_rejects_orthogonal_direction_magnitude_as_a_signal():
weight = torch.tensor([[4.0, 3.0], [2.0, -1.0]], dtype=torch.float64)
unit = project_weight_against_direction(weight, torch.tensor([1.0, 0.0], dtype=torch.float64))
scaled = project_weight_against_direction(weight, torch.tensor([5.0, 0.0], dtype=torch.float64))
assert torch.allclose(unit.weight, scaled.weight, atol=1e-12)
assert unit.layout == scaled.layout == "standard"
def test_row_fraction_selects_largest_coefficients_and_is_permutation_equivariant():
coeff = torch.tensor([[0.5], [-3.0], [2.0], [0.1]])
selected = select_projection_coefficients(coeff, 0.5)
assert selected.tolist() == [[0.0], [-3.0], [2.0], [0.0]]
permutation = torch.tensor([2, 0, 3, 1])
permuted = select_projection_coefficients(coeff[permutation], 0.5)
assert torch.allclose(permuted, selected[permutation])
def test_projection_row_fraction_removes_only_selected_rows():
weight = torch.tensor([[10.0, 1.0], [1.0, 7.0], [-5.0, 2.0], [0.2, 9.0]])
direction = torch.tensor([1.0, 0.0])
projected = project_weight_against_direction(weight, direction, projection_row_fraction=0.5)
assert torch.allclose(projected.weight[:, 0], torch.tensor([0.0, 1.0, 0.0, 0.2]))
assert torch.allclose(projected.weight[:, 1], weight[:, 1])
def test_projection_row_fraction_keeps_only_the_two_largest_magnitudes():
coeff = torch.tensor([[0.5], [-3.0], [2.0], [0.1]], dtype=torch.float64)
selected = select_projection_coefficients(coeff, 0.5)
assert torch.equal(selected != 0, torch.tensor([[False], [True], [True], [False]]))
assert torch.allclose(selected.abs().sum(), torch.tensor(5.0, dtype=torch.float64))
@pytest.mark.parametrize("regularization", [-1.0, 0.25, 1.25])
def test_finite_regularization_values_are_applied_without_unit_interval_clamping(regularization):
weight = torch.tensor([[4.0, 3.0]], dtype=torch.float64)
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
scale = 1.0 - regularization
projected = project_weight_against_direction(weight, direction, regularization=regularization)
assert projected.projected
assert torch.allclose(projected.weight, _reference_projection(weight, direction, scale), atol=1e-12)
def test_harmless_pc_removal_orthogonalizes_against_dominant_component():
subspace = torch.tensor([[1.0, 1.0, 0.0], [0.5, 0.0, 1.0]], dtype=torch.float64)
harmless = torch.tensor(
[[-2.0, 0.0, 0.0], [0.0, 0.0, 0.0], [2.0, 0.0, 0.0], [4.0, 0.0, 0.0]],
dtype=torch.float64,
)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual[:, 0], torch.zeros(2, dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual.norm(dim=-1), torch.ones(2, dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual @ residual.T, torch.eye(2, dtype=torch.float64), atol=1e-12)
@pytest.mark.parametrize(
("subspace", "harmless", "pc_count"),
[
(
torch.tensor([[1.0, 0.0]], dtype=torch.float64),
torch.eye(3, 2, dtype=torch.float64),
0,
),
(
torch.tensor([[1.0, 0.0]], dtype=torch.float64),
torch.eye(2, dtype=torch.float64),
1,
),
(
torch.empty((0, 2), dtype=torch.float64),
torch.eye(3, 2, dtype=torch.float64),
1,
),
],
)
def test_harmless_pc_removal_noops_when_preconditions_are_not_met(
subspace,
harmless,
pc_count,
):
residual = remove_harmless_principal_components(subspace, harmless, pc_count)
assert residual is subspace
def test_harmless_pc_removal_returns_subspace_when_svd_fails(monkeypatch):
subspace = torch.tensor([[1.0, 0.5]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
def fail_svd(*_args, **_kwargs):
raise RuntimeError("svd fixture failure")
monkeypatch.setattr(torch.linalg, "svd", fail_svd)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert residual is subspace
def test_harmless_pc_removal_noops_when_svd_returns_no_components():
subspace = torch.tensor([[1.0]], dtype=torch.float64)
harmless = torch.empty((3, 0), dtype=torch.float64)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert residual is subspace
def test_harmless_pc_removal_subtracts_pc_component_before_row_normalization():
subspace = torch.tensor([[1.0, 0.5]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual, torch.tensor([[0.0, 1.0]], dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual.norm(dim=-1), torch.ones(1, dtype=torch.float64), atol=1e-12)
def test_harmless_pc_removal_treats_exact_epsilon_residual_as_usable_signal():
subspace = torch.tensor([[1.0, 1e-8]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual, torch.tensor([[0.0, 1.0]], dtype=torch.float64), atol=1e-12)
def test_harmless_pc_removal_preserves_subunit_residuals_instead_of_restoring_original_row():
subspace = torch.tensor([[1.0, 0.5]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual[:, 0], torch.zeros(1, dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual[:, 1], torch.ones(1, dtype=torch.float64), atol=1e-12)
def test_harmless_pc_removal_restores_only_rows_with_near_zero_residuals():
subspace = torch.tensor([[1.0, 0.0], [1.0, 0.5]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual @ residual.T, torch.eye(2, dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual[0], torch.tensor([1.0, 0.0], dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual[1], torch.tensor([0.0, 1.0], dtype=torch.float64), atol=1e-12)
def test_harmless_pc_removal_near_zero_fallback_is_per_row_before_qr():
subspace = torch.tensor([[1.0, 1e-9, 0.0], [0.2, 0.5, 1.0]], dtype=torch.float64)
harmless = torch.tensor(
[[-2.0, 0.0, 0.0], [0.0, 0.0, 0.0], [2.0, 0.0, 0.0]],
dtype=torch.float64,
)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert torch.allclose(residual[0], subspace[0], atol=1e-12)
assert torch.allclose(residual @ residual.T, torch.eye(2, dtype=torch.float64), atol=1e-8)
def test_harmless_pc_removal_replays_deterministically_for_singular_inputs():
subspace = torch.tensor([[0.0, 1.0, 1.0], [0.0, 2.0, 2.0]], dtype=torch.float64)
harmless = torch.ones((4, 3), dtype=torch.float64)
first = remove_harmless_principal_components(subspace, harmless, pc_count=2)
second = remove_harmless_principal_components(subspace, harmless, pc_count=2)
assert torch.allclose(first, second, atol=0.0, rtol=0.0)
assert torch.isfinite(first).all()
def test_shield_atom_residualization_handles_rank_deficient_atoms():
subspace = torch.tensor([[1.0, 1.0, 0.0], [1.0, 0.0, 1.0]], dtype=torch.float64)
atoms = torch.tensor([[1.0, 0.0, 0.0], [2.0, 0.0, 0.0]], dtype=torch.float64)
residual = residualize_against_shield_atoms(subspace, atoms, ridge=1e-6)
assert torch.allclose(residual[:, 0], torch.zeros(2, dtype=torch.float64), atol=2e-6)
assert torch.allclose(residual.norm(dim=-1), torch.ones(2, dtype=torch.float64), atol=1e-12)
assert torch.allclose(residual @ residual.T, torch.eye(2, dtype=torch.float64), atol=1e-12)
def test_shield_atom_residualization_noops_when_atoms_or_subspace_are_empty():
subspace = torch.tensor([[1.0, 0.0]], dtype=torch.float64)
atoms = torch.empty((0, 2), dtype=torch.float64)
empty_subspace = torch.empty((0, 2), dtype=torch.float64)
assert residualize_against_shield_atoms(subspace, atoms, ridge=1e-3) is subspace
assert residualize_against_shield_atoms(empty_subspace, torch.eye(2), ridge=1e-3) is empty_subspace
def test_shield_atom_residualization_returns_subspace_when_solve_fails(monkeypatch):
subspace = torch.tensor([[1.0, 1.0]], dtype=torch.float64)
atoms = torch.tensor([[1.0, 0.0]], dtype=torch.float64)
def fail_solve(*_args, **_kwargs):
raise RuntimeError("solve fixture failure")
monkeypatch.setattr(torch.linalg, "solve", fail_solve)
residual = residualize_against_shield_atoms(subspace, atoms, ridge=1e-6)
assert residual is subspace
def test_shield_atom_residualization_uses_mixed_dtype_compute_but_returns_subspace_dtype():
subspace = torch.tensor([[1.0, 1.0, 0.0]], dtype=torch.float32)
atoms = torch.tensor([[1.0, 0.0, 0.0]], dtype=torch.float64)
residual = residualize_against_shield_atoms(subspace, atoms, ridge=1e-6)
assert residual.dtype == torch.float32
assert residual.device == subspace.device
assert torch.allclose(residual, torch.tensor([[0.0, 1.0, 0.0]], dtype=torch.float32), atol=2e-6)
def test_shield_atom_residualization_upcasts_atoms_to_match_float64_subspace_compute():
subspace = torch.tensor([[1.0, 1.0, 0.0]], dtype=torch.float64)
atoms = torch.tensor([[1.0, 0.0, 0.0]], dtype=torch.float32)
residual = residualize_against_shield_atoms(subspace, atoms, ridge=1e-6)
assert residual.dtype == torch.float64
assert torch.allclose(residual, torch.tensor([[0.0, 1.0, 0.0]], dtype=torch.float64), atol=2e-6)
def test_shield_atom_residualization_keeps_float64_precision_with_float32_atoms():
subspace = torch.tensor([[1.0, 1e-4, 1.0 - 1e-4]], dtype=torch.float64)
atoms = torch.tensor([[1.0, 1e-4, 0.0], [1e-4, 1.0, 1e-4]], dtype=torch.float32)
residual = residualize_against_shield_atoms(subspace, atoms, ridge=1e-12)
assert residual.dtype == torch.float64
assert residual[0, 0] > 1e-8
assert torch.allclose(
residual,
torch.tensor(
[[1.0000999534067183e-08, -9.9999997973787514e-05, 0.9999999950000001]],
dtype=torch.float64,
),
atol=1e-15,
)
def test_harmless_pc_removal_upcasts_half_precision_for_svd_then_returns_input_dtype():
subspace = torch.tensor([[1.0, 0.5]], dtype=torch.float16)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float16)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert residual.dtype == torch.float16
assert torch.allclose(residual.float(), torch.tensor([[0.0, 1.0]]), atol=1e-3)
def test_dtype_and_device_are_preserved_for_projection_and_residualization():
weight = torch.tensor([[1.0, 2.0]], dtype=torch.float32)
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
subspace = torch.tensor([[1.0, 1.0]], dtype=torch.float32)
harmless = torch.tensor([[-1.0, 0.0], [0.0, 0.0], [1.0, 0.0]], dtype=torch.float64)
projected = project_weight_against_direction(weight, direction)
residual = remove_harmless_principal_components(subspace, harmless, pc_count=1)
assert projected.weight.dtype == torch.float32
assert projected.weight.device == weight.device
assert residual.dtype == torch.float32
assert residual.device == subspace.device
@pytest.mark.parametrize("bad_value", [float("nan"), float("inf")])
def test_projection_non_finite_policy_is_skip_without_mutation(bad_value):
weight = torch.tensor([[1.0, 2.0], [bad_value, 4.0]])
direction = torch.tensor([1.0, 0.0])
projected = project_weight_against_direction(weight, direction)
assert not projected.projected
assert torch.allclose(projected.weight, weight, equal_nan=True)
def test_zero_inputs_follow_existing_fallback_policy_without_non_finite_output():
zero_subspace = torch.zeros((2, 3), dtype=torch.float64)
harmless = torch.tensor([[-1.0, 0.0, 0.0], [0.0, 0.0, 0.0], [1.0, 0.0, 0.0]])
atoms = torch.tensor([[1.0, 0.0, 0.0]])
pc_residual = remove_harmless_principal_components(zero_subspace, harmless, pc_count=1)
shield_residual = residualize_against_shield_atoms(zero_subspace, atoms, ridge=1e-3)
assert torch.isfinite(pc_residual).all()
assert torch.isfinite(shield_residual).all()
assert torch.allclose(pc_residual, zero_subspace)
assert torch.allclose(shield_residual, zero_subspace)
def test_regularized_projection_replay_matches_closed_form_decay():
weight = torch.tensor([[4.0, 3.0]], dtype=torch.float64)
direction = torch.tensor([1.0, 0.0], dtype=torch.float64)
regularization = 0.25
scale = 1.0 - regularization
first = project_weight_against_direction(weight, direction, regularization=regularization)
second = project_weight_against_direction(first.weight, direction, regularization=regularization)
assert torch.allclose(first.weight[:, 0], weight[:, 0] * regularization)
assert torch.allclose(second.weight[:, 0], weight[:, 0] * math.pow(regularization, 2))
assert torch.allclose(first.weight, _reference_projection(weight, direction, scale))
def test_abliteration_pipeline_math_wrappers_delegate_to_contract_helpers():
subspace = torch.tensor([[1.0, 0.0], [1.0, 1.0]], dtype=torch.float64)
harmless = torch.tensor([[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0]], dtype=torch.float64)
atoms = torch.tensor([[1.0, 0.0]], dtype=torch.float64)
coeff = torch.tensor([[0.1], [3.0], [-2.0]], dtype=torch.float64)
assert torch.allclose(
AbliterationPipeline._orthogonalize_subspace(subspace),
orthogonalize_subspace_rows(subspace),
)
assert torch.allclose(
AbliterationPipeline(None)._remove_harmless_principal_components(
subspace,
harmless,
1,
),
remove_harmless_principal_components(subspace, harmless, 1),
)
assert torch.allclose(
AbliterationPipeline(None)._residualize_against_shield_atoms(subspace, atoms, 1e-6),
residualize_against_shield_atoms(subspace, atoms, 1e-6),
)
assert torch.equal(
AbliterationPipeline._select_projection_coefficients(coeff, 0.5),
select_projection_coefficients(coeff, 0.5),
)
def test_project_out_advanced_replaces_quantized_weight_after_successful_projection(monkeypatch):
linear = torch.nn.Linear(2, 2, bias=False)
with torch.no_grad():
linear.weight.copy_(torch.tensor([[2.0, 0.0], [0.0, 2.0]]))
module = SimpleNamespace(o_proj=linear)
replacement_calls = []
monkeypatch.setattr(
AbliterationPipeline,
"_dequantize_weight",
staticmethod(lambda proj: (proj.weight.data, True)),
)
monkeypatch.setattr(
AbliterationPipeline,
"_replace_quantized_weight",
staticmethod(lambda proj, weight: replacement_calls.append((proj, weight.clone()))),
)
count = AbliterationPipeline._project_out_advanced(
module,
torch.tensor([1.0, 0.0]),
["o_proj"],
)
assert count == 1
assert len(replacement_calls) == 1
assert replacement_calls[0][0] is linear
assert torch.allclose(replacement_calls[0][1][:, 0], torch.zeros(2))
def test_quantized_parameter_wrapper_uses_declared_markers_only():
params4bit = type("Params4bit", (), {})()
ordinary = SimpleNamespace()
marked = SimpleNamespace(quant_state=object())
assert AbliterationPipeline._is_quantized_param(params4bit) is True
assert AbliterationPipeline._is_quantized_param(marked) is True
assert AbliterationPipeline._is_quantized_param(ordinary) is False
def test_dequantize_weight_returns_float_storage_without_copying():
linear = torch.nn.Linear(2, 2, bias=False)
weight, requires_replacement = AbliterationPipeline._dequantize_weight(linear)
assert weight.data_ptr() == linear.weight.data.data_ptr()
assert requires_replacement is False
def test_dequantize_weight_promotes_integer_storage_for_safe_projection():
module = SimpleNamespace(
weight=torch.nn.Parameter(
torch.tensor([[1, 2], [3, 4]], dtype=torch.uint8),
requires_grad=False,
),
)
weight, requires_replacement = AbliterationPipeline._dequantize_weight(module)
assert weight.dtype is torch.float32
assert torch.equal(weight, torch.tensor([[1.0, 2.0], [3.0, 4.0]]))
assert requires_replacement is True
def test_packed_module_dequantization_is_cloned_before_projection():
packed_type = type("QuantLinear", (), {})
packed = packed_type()
source = torch.tensor([[1.0, 2.0]])
packed.dequantize = lambda: source
weight, requires_replacement = AbliterationPipeline._dequantize_weight(packed)
assert torch.equal(weight, source)
assert weight.data_ptr() != source.data_ptr()
assert requires_replacement is True
def test_packed_module_without_safe_dequantizer_fails_closed():
packed_type = type("WQLinear_GEMM", (), {})
packed = packed_type()
packed.qweight = torch.ones(1, dtype=torch.int32)
packed.scales = torch.ones(1)
with pytest.raises(RuntimeError, match=r"no dequantize\(\) method available"):
AbliterationPipeline._dequantize_weight(packed)
def test_weight_replacement_preserves_float_parameter_contract():
linear = torch.nn.Linear(2, 2, bias=False)
modified = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
AbliterationPipeline._replace_quantized_weight(linear, modified)
assert torch.equal(linear.weight, modified)
assert linear.weight.requires_grad is True
def test_packed_weight_replacement_uses_the_module_repacker():
packed_type = type("QuantLinear", (), {})
packed = packed_type()
packed.qweight = torch.ones(1)
packed.scales = torch.tensor([0.5])
calls = []
packed.pack = lambda weight, scales: calls.append((weight.clone(), scales.clone()))
modified = torch.tensor([[1.0, 2.0]])
AbliterationPipeline._replace_quantized_weight(packed, modified)
assert len(calls) == 1
assert torch.equal(calls[0][0], modified)
assert torch.equal(calls[0][1], packed.scales)