etomoscow/mff_lora / code /tests /test_init_shapes.py
etomoscow's picture
download
raw
2.81 kB
"""Shape checks for all init functions."""
from __future__ import annotations
import pytest
import torch
from mfflora.init import (
eva_init,
lora_ga_init,
milora_init,
mff_lora_init,
pissa_init,
random_init,
filet_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, n, m
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)
@pytest.mark.parametrize("rank", [2, 4, 8])
@pytest.mark.parametrize("selection", ["top", "bottom", "energy"])
@pytest.mark.parametrize("strategy", ["alpha", "beta", "gamma"])
def test_mff_lora_shapes(small_layer, rank, selection, strategy):
W, n, m = small_layer
A = _psd(n, seed=10)
B = _psd(m, seed=11)
res = mff_lora_init(W, A, B, rank, selection=selection, strategy=strategy)
assert res.U.shape == (n, rank)
assert res.V.shape == (rank, m)
if strategy == "gamma":
assert res.residual is not None
assert res.residual.shape == (n, m)
else:
assert res.residual is None
@pytest.mark.parametrize("rank", [2, 4, 8])
def test_random_shapes(small_layer, rank):
W, n, m = small_layer
U, V = random_init(W, rank)
assert U.shape == (n, rank)
assert V.shape == (rank, m)
assert torch.equal(U, torch.zeros_like(U))
@pytest.mark.parametrize("rank", [2, 4])
def test_pissa_shapes(small_layer, rank):
W, n, m = small_layer
res = pissa_init(W, rank)
assert res.U.shape == (n, rank)
assert res.V.shape == (rank, m)
assert res.residual.shape == (n, m)
@pytest.mark.parametrize("rank", [2, 4])
def test_milora_shapes(small_layer, rank):
W, n, m = small_layer
res = milora_init(W, rank)
assert res.U.shape == (n, rank)
assert res.V.shape == (rank, m)
assert res.residual.shape == (n, m)
@pytest.mark.parametrize("rank", [2, 4])
def test_lora_ga_shapes(small_layer, rank):
W, n, m = small_layer
g = torch.Generator().manual_seed(7)
fake_grad = torch.randn(n, m, generator=g)
U, V = lora_ga_init(W, rank, gradient=fake_grad)
assert U.shape == (n, rank)
assert V.shape == (rank, m)
@pytest.mark.parametrize("rank", [2, 4])
def test_eva_shapes(small_layer, rank):
W, n, m = small_layer
cov = _psd(m, seed=12)
U, V = eva_init(W, rank, activation_cov=cov)
assert U.shape == (n, rank)
assert V.shape == (rank, m)
@pytest.mark.parametrize("rank", [2, 4])
def test_filet_shapes(small_layer, rank):
W, n, m = small_layer
sx = _psd(m, seed=13)
sy = _psd(n, seed=14)
U, V = filet_init(W, rank, sx, sy)
assert U.shape == (n, rank)
assert V.shape == (rank, m)

Xet Storage Details

Size:
2.81 kB
·
Xet hash:
26e928356b5622718ab35c9fdb9f40807ee0044c2915b6f062a229eb869b885d

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