etomoscow/mff_lora / code /tests /test_init_zero_at_start.py
etomoscow's picture
download
raw
3.13 kB
"""ΔW = 0 at init invariant for strategies α and β; PiSSA preservation for γ."""
from __future__ import annotations
import pytest
import torch
from mfflora.init import milora_init, mff_lora_init, pissa_init, random_init
@pytest.fixture
def small_layer():
n, m = 24, 16
g = torch.Generator().manual_seed(0)
W = torch.randn(n, m, generator=g) * 0.1
return W
def _psd(d: int, seed: int) -> torch.Tensor:
g = torch.Generator().manual_seed(seed)
L = torch.randn(d, d, generator=g)
return L @ L.T + 1e-3 * torch.eye(d)
def test_random_init_zero_delta(small_layer):
U, V = random_init(small_layer, rank=4)
assert torch.equal(U @ V, torch.zeros_like(small_layer))
@pytest.mark.parametrize("selection", ["top", "bottom", "energy"])
def test_mff_strategy_alpha_zero_delta(small_layer, selection):
n, m = small_layer.shape
A = _psd(n, seed=10)
B = _psd(m, seed=11)
res = mff_lora_init(small_layer, A, B, rank=4, selection=selection, strategy="alpha")
assert torch.allclose(res.U @ res.V, torch.zeros_like(small_layer), atol=1e-6)
@pytest.mark.parametrize("selection", ["top", "bottom", "energy"])
def test_mff_strategy_beta_zero_delta(small_layer, selection):
n, m = small_layer.shape
A = _psd(n, seed=10)
B = _psd(m, seed=11)
res = mff_lora_init(small_layer, A, B, rank=4, selection=selection, strategy="beta")
assert torch.allclose(res.U @ res.V, torch.zeros_like(small_layer), atol=1e-6)
def test_mff_strategy_gamma_preserves_forward(small_layer):
n, m = small_layer.shape
A = _psd(n, seed=10)
B = _psd(m, seed=11)
res = mff_lora_init(small_layer, A, B, rank=4, selection="top", strategy="gamma")
assert res.residual is not None
g = torch.Generator().manual_seed(99)
x = torch.randn(7, m, generator=g)
y_orig = x @ small_layer.T
y_lora = x @ (res.residual + res.U @ res.V).T
assert torch.allclose(y_orig, y_lora, atol=1e-4)
def test_pissa_preserves_forward(small_layer):
res = pissa_init(small_layer, rank=4)
g = torch.Generator().manual_seed(100)
x = torch.randn(7, small_layer.shape[1], generator=g)
y_orig = x @ small_layer.T
y_lora = x @ (res.residual + res.U @ res.V).T
assert torch.allclose(y_orig, y_lora, atol=1e-4)
def test_milora_preserves_forward(small_layer):
res = milora_init(small_layer, rank=4)
g = torch.Generator().manual_seed(101)
x = torch.randn(7, small_layer.shape[1], generator=g)
y_orig = x @ small_layer.T
y_lora = x @ (res.residual + res.U @ res.V).T
assert torch.allclose(y_orig, y_lora, atol=1e-4)
def test_top_picks_higher_eigvals_than_bottom(small_layer):
n, m = small_layer.shape
A = _psd(n, seed=10)
B = _psd(m, seed=11)
top = mff_lora_init(small_layer, A, B, rank=4, selection="top", strategy="alpha")
bot = mff_lora_init(small_layer, A, B, rank=4, selection="bottom", strategy="alpha")
# Quadratic form on A is higher for top columns than bottom columns.
top_quad = ((A @ top.U) * top.U).sum(0).mean()
bot_quad = ((A @ bot.U) * bot.U).sum(0).mean()
assert top_quad > bot_quad

Xet Storage Details

Size:
3.13 kB
·
Xet hash:
f6148d92f7cda5b313c3988c649345227d6d2b93ca51fe2c4ef88099d3af7aa6

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.