| """Tests for src/mfflora/factors.py.""" | |
| from __future__ import annotations | |
| import pytest | |
| import torch | |
| from mfflora.factors import ( | |
| fisher_energy_score, | |
| load_factors, | |
| topk_eigvecs, | |
| ) | |
| def test_topk_top_recovers_known_directions(random_psd_factory): | |
| n = 50 | |
| M = random_psd_factory(n, seed=1) | |
| vecs, vals = topk_eigvecs(M, k=5, ascending=False, method="exact") | |
| assert vecs.shape == (n, 5) | |
| assert vals.shape == (5,) | |
| # Eigenvalues descending, all positive (PSD). | |
| assert torch.all(vals[:-1] >= vals[1:] - 1e-5) | |
| assert vals.min() > 0 | |
| # Ground-truth comparison. | |
| eigvals_true, _ = torch.linalg.eigh(M) | |
| top_true = eigvals_true.flip(0)[:5] | |
| assert torch.allclose(vals, top_true, atol=1e-4) | |
| def test_topk_bottom_returns_smallest(random_psd_factory): | |
| n = 50 | |
| M = random_psd_factory(n, seed=2) | |
| vecs, vals = topk_eigvecs(M, k=5, ascending=True, method="exact") | |
| assert vecs.shape == (n, 5) | |
| eigvals_true, _ = torch.linalg.eigh(M) | |
| bottom_true = eigvals_true[:5] | |
| assert torch.allclose(vals, bottom_true, atol=1e-4) | |
| def test_topk_reconstruction(random_psd_factory): | |
| n = 80 | |
| k = 60 | |
| M = random_psd_factory(n, seed=3) | |
| vecs, vals = topk_eigvecs(M, k=k, ascending=False, method="exact") | |
| # vecs orthonormal columns. | |
| inner = vecs.T @ vecs | |
| assert torch.allclose(inner, torch.eye(k), atol=1e-4) | |
| # Best rank-k Frobenius approximation: ||M - V Λ V^T||_F bounded by sum of | |
| # remaining eigenvalues. | |
| M_approx = vecs @ torch.diag(vals) @ vecs.T | |
| eigvals_true, _ = torch.linalg.eigh(M) | |
| tail = eigvals_true[: n - k] | |
| assert (M - M_approx).norm() <= tail.norm() + 1e-3 | |
| def test_randomized_topk_close_to_exact(random_psd_factory): | |
| """For a Wishart matrix the spectrum has no clear gap; randomized | |
| subspace iteration converges only to a few-percent relative error in a | |
| handful of iterations. We use it for *direction* recovery (subspace | |
| overlap), not exact eigenvalue agreement.""" | |
| n = 200 | |
| k = 8 | |
| M = random_psd_factory(n, seed=4) | |
| vecs_e, vals_e = topk_eigvecs(M, k=k, method="exact") | |
| vecs_r, vals_r = topk_eigvecs(M, k=k, method="randomized", n_iter=10, oversampling=20) | |
| # Eigenvalue agreement loose: subspace iteration biases low for trailing | |
| # ranks; tolerate 5% relative error. | |
| rel = (vals_e - vals_r).abs() / vals_e.abs() | |
| assert float(rel.max()) < 0.05 | |
| # Subspace overlap close to 1. | |
| cos = torch.linalg.svdvals(vecs_e.T @ vecs_r).clamp(0.0, 1.0) | |
| assert float(cos.min()) > 0.9 | |
| def test_fisher_energy_score_matches_quadratic_form(random_psd_factory): | |
| n, m, k = 30, 20, 5 | |
| A = random_psd_factory(n, seed=5) | |
| B = random_psd_factory(m, seed=6) | |
| Va, _ = topk_eigvecs(A, k) | |
| Vb, _ = topk_eigvecs(B, k) | |
| energy = fisher_energy_score(Va, Vb, A, B) | |
| # Manual computation. | |
| expected = torch.tensor( | |
| [(Va[:, i] @ A @ Va[:, i]) * (Vb[:, i] @ B @ Vb[:, i]) for i in range(k)] | |
| ) | |
| assert torch.allclose(energy.float(), expected.float(), atol=1e-4) | |
| def test_load_real_factor_shape(factor_library_root, has_factor_library): | |
| if not has_factor_library: | |
| pytest.skip("Factor library not available") | |
| factors = load_factors("en", layer=0, module="gate_proj", factor_dir=factor_library_root) | |
| assert factors.A.shape == (14336, 14336) | |
| assert factors.B.shape == (4096, 4096) | |
| assert factors.A.dtype == torch.float32 | |
| assert factors.layer_name == "model.layers.0.mlp.gate_proj" | |
Xet Storage Details
- Size:
- 3.5 kB
- Xet hash:
- 81d2a80ca80f7e08b2838f516b496d1163d10d8137aae8477b883910d9d15648
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.