etomoscow/mff_lora / code /tests /test_factors.py
etomoscow's picture
download
raw
3.5 kB
"""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.