Agnes-3.0-Flash-GGUF / build /test_projection.py
0xKitkat's picture
Upload provenance, runtime files and measured evaluations
35514b1 verified
Raw History Blame
1.25 kB
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()