File size: 1,051 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
import copy
import torch
from nexora.adapters import LoRALinear, inject_lora, merge_lora
from nexora.model import ModelConfig, NexoraLM


def test_lora_freeze_and_merge():
    torch.manual_seed(2)
    base = torch.nn.Linear(8, 12)
    layer = LoRALinear(base, rank=2)
    x = torch.randn(3, 8)
    torch.testing.assert_close(layer(x), base(x))
    opt = torch.optim.SGD([layer.a, layer.b], lr=.1)
    old = base.weight.detach().clone()
    layer(x).square().mean().backward()
    opt.step()
    assert torch.equal(base.weight, old)
    assert layer.b.abs().sum() > 0
    torch.testing.assert_close(layer(x), layer.merged()(x))


def test_inject_only_adapter_trainable():
    m = NexoraLM(ModelConfig(hidden_size=32, layers=1, heads=4, kv_heads=2, intermediate_size=64))
    assert inject_lora(m) == ["blocks.0.attn.q", "blocks.0.attn.v"]
    assert all(name.endswith((".a", ".b")) for name, p in m.named_parameters() if p.requires_grad)
    merged = merge_lora(copy.deepcopy(m))
    assert not any(isinstance(x, LoRALinear) for x in merged.modules())