import torch from abliterate import project def test_output_projection_removes_only_selected_component(): torch.manual_seed(19) matrix = torch.randn(7, 31, dtype=torch.float32) direction = torch.randn(7) direction /= direction.norm() edited = project(matrix, direction, 1.0) torch.testing.assert_close(direction @ edited, torch.zeros(31), atol=2e-6, rtol=0) perpendicular = torch.randn(7) perpendicular -= direction * (perpendicular @ direction) torch.testing.assert_close(perpendicular @ matrix, perpendicular @ edited, atol=3e-6, rtol=1e-5) def test_embedding_projection_and_partial_strength(): torch.manual_seed(20) matrix = torch.randn(2053, 7) direction = torch.randn(7) direction /= direction.norm() edited = project(matrix, direction, 0.7, embedding=True) torch.testing.assert_close(edited @ direction, 0.3 * (matrix @ direction), atol=2e-6, rtol=1e-4) def test_bf16_is_finite_and_preserves_shape(): matrix = torch.randn(7, 2051).bfloat16() direction = torch.randn(7) direction /= direction.norm() edited = project(matrix, direction, 1.0) assert edited.shape == matrix.shape assert edited.dtype == torch.bfloat16 assert torch.isfinite(edited).all()