| """Δ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 | |
| 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)) | |
| 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) | |
| 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.