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