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