"""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