File size: 2,248 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
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
"""Small transparent LoRA implementation for the reference model's linear layers.

For external production models use PEFT; this is an inspectable training experiment.
"""
import math
import torch
from torch import nn


class LoRALinear(nn.Module):
    def __init__(self, base: nn.Linear, rank=4, alpha=8):
        super().__init__()
        if rank < 1 or alpha <= 0:
            raise ValueError("Invalid LoRA hyperparameters")
        self.base, self.scale = base, alpha/rank
        self.base.requires_grad_(False)
        self.a = nn.Parameter(torch.empty(rank, base.in_features, device=base.weight.device, dtype=base.weight.dtype))
        self.b = nn.Parameter(torch.zeros(base.out_features, rank, device=base.weight.device, dtype=base.weight.dtype))
        nn.init.kaiming_uniform_(self.a, a=math.sqrt(5))

    def forward(self, x):
        return self.base(x) + (x @ self.a.T @ self.b.T)*self.scale

    def merged(self):
        base = nn.Linear(self.base.in_features, self.base.out_features, bias=self.base.bias is not None,
                         device=self.base.weight.device, dtype=self.base.weight.dtype)
        with torch.no_grad():
            base.weight.copy_(self.base.weight + (self.b @ self.a)*self.scale)
            if base.bias is not None:
                base.bias.copy_(self.base.bias)
        return base


def inject_lora(model, targets=("q", "v"), rank=4, alpha=8):
    model.requires_grad_(False)
    replaced = []
    for name, module in list(model.named_modules()):
        if isinstance(module, nn.Linear) and name.rsplit(".", 1)[-1] in targets:
            parent_name, _, leaf = name.rpartition(".")
            parent = model.get_submodule(parent_name) if parent_name else model
            setattr(parent, leaf, LoRALinear(module, rank, alpha))
            replaced.append(name)
    if not replaced:
        raise ValueError("No target linear modules found")
    return replaced


def merge_lora(model):
    for name, module in list(model.named_modules()):
        if isinstance(module, LoRALinear):
            parent_name, _, leaf = name.rpartition(".")
            parent = model.get_submodule(parent_name) if parent_name else model
            setattr(parent, leaf, module.merged())
    return model