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
|