| """Tests for Fisher-informed LoRA initialization strategies.""" | |
| import pytest | |
| import torch | |
| from mfflora.init.fisher_whitened import ( | |
| fisher_peft_init, | |
| fisher_peft_low_init, | |
| fisher_peft_bi_init, | |
| fisher_orthogonal_init, | |
| fws_safe_init, | |
| ) | |
| def _make_factors(n, m, rank, seed=42): | |
| """Create synthetic W, SX, SY with known structure.""" | |
| torch.manual_seed(seed) | |
| W = torch.randn(n, m) * 0.1 | |
| V_x = torch.linalg.qr(torch.randn(m, m)).Q | |
| lam_x = torch.randn(m).abs() | |
| lam_x[:rank] *= 10 | |
| SX = V_x @ torch.diag(lam_x) @ V_x.T | |
| SX = (SX + SX.T) / 2 | |
| V_y = torch.linalg.qr(torch.randn(n, n)).Q | |
| lam_y = torch.randn(n).abs() | |
| lam_y[:rank] *= 10 | |
| SY = V_y @ torch.diag(lam_y) @ V_y.T | |
| SY = (SY + SY.T) / 2 | |
| return W, SX, SY | |
| class TestFisherPEFT: | |
| def test_shapes(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_peft_init(W, SX, r) | |
| assert B.shape == (n, r) | |
| assert A.shape == (r, m) | |
| def test_delta_is_zero(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_peft_init(W, SX, r) | |
| assert (B @ A).norm() == pytest.approx(0.0, abs=1e-12) | |
| def test_a_is_orthonormal(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| _, A = fisher_peft_init(W, SX, r) | |
| assert torch.allclose(A @ A.T, torch.eye(r), atol=1e-4) | |
| class TestFisherPEFTLow: | |
| def test_shapes(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_peft_low_init(W, SX, r) | |
| assert B.shape == (n, r) | |
| assert A.shape == (r, m) | |
| def test_delta_is_zero(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_peft_low_init(W, SX, r) | |
| assert (B @ A).norm() == pytest.approx(0.0, abs=1e-12) | |
| class TestFisherPEFTBi: | |
| def test_shapes_high(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| B, A = fisher_peft_bi_init(W, SX, SY, r, select="high") | |
| assert B.shape == (n, r) | |
| assert A.shape == (r, m) | |
| def test_shapes_low(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| B, A = fisher_peft_bi_init(W, SX, SY, r, select="low") | |
| assert B.shape == (n, r) | |
| assert A.shape == (r, m) | |
| def test_delta_is_zero(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| B, A = fisher_peft_bi_init(W, SX, SY, r, select="high") | |
| assert (B @ A).norm() == pytest.approx(0.0, abs=1e-12) | |
| def test_high_vs_low_differ(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| _, A_high = fisher_peft_bi_init(W, SX, SY, r, select="high") | |
| _, A_low = fisher_peft_bi_init(W, SX, SY, r, select="low") | |
| assert not torch.allclose(A_high, A_low) | |
| class TestFisherOrthogonal: | |
| def test_shapes(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_orthogonal_init(W, SX, r) | |
| assert B.shape == (n, r) | |
| assert A.shape == (r, m) | |
| def test_delta_is_small(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| B, A = fisher_orthogonal_init(W, SX, r) | |
| delta = B @ A | |
| assert delta.norm() < B.norm() * 1.01 | |
| def test_a_is_orthonormal(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, _ = _make_factors(n, m, r) | |
| _, A = fisher_orthogonal_init(W, SX, r) | |
| assert torch.allclose(A @ A.T, torch.eye(r), atol=1e-4) | |
| class TestFWSSafe: | |
| def test_shapes(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| U, V = fws_safe_init(W, SX, SY, r) | |
| assert U.shape == (n, r) | |
| assert V.shape == (r, m) | |
| def test_magnitude_control(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| U, V = fws_safe_init(W, SX, SY, r, init_scale=0.001) | |
| ratio = (U @ V).norm() / W.norm() | |
| assert ratio == pytest.approx(0.001, rel=0.01) | |
| def test_delta_is_safe(self): | |
| n, m, r = 64, 32, 8 | |
| W, SX, SY = _make_factors(n, m, r) | |
| U, V = fws_safe_init(W, SX, SY, r, init_scale=0.001) | |
| ratio = (U @ V).norm() / W.norm() | |
| assert ratio < 0.002 # within safe zone | |
Xet Storage Details
- Size:
- 4.34 kB
- Xet hash:
- 806384e32066b18a57e35741d9b365895327ebfbbc6ed498c1964dd3d0727a0b
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.