NEXORA / nexora /adapters.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
2.25 kB
"""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