"""LoRA: teach the chat model a new style by training two thin matrices per layer. python lora.py --base runs/chat/ckpt_final.pt --chats sft/style/pirate.jsonl --out runs/pirate A frozen weight W gets a detour: y = W x + (alpha / r) · B (A x), A: r × in, B: out × r. B starts at zero, so at step 0 the model is exactly the chat model. Only A and B learn — with rank 8 that is about 2.4 % of Sprout's parameters. At the end B·A is merged into W, and the result is an ordinary Sprout checkpoint again. """ import argparse import json import math import os import random import time import torch import torch.nn as nn from model import GPT, Config from sft import pack, render from tokenizer import Tokenizer from train import get_device, lr_factor class LoRALinear(nn.Module): def __init__(self, base, rank=8, alpha=16): super().__init__() self.base = base self.A = nn.Parameter(torch.randn(rank, base.in_features) / math.sqrt(base.in_features)) self.B = nn.Parameter(torch.zeros(base.out_features, rank)) self.scale = alpha / rank def forward(self, x): return self.base(x) + (x @ self.A.T @ self.B.T) * self.scale def merge(self): self.base.weight.data += (self.B @ self.A).to(self.base.weight.dtype) * self.scale return self.base TARGETS = [('attn', 'qkv'), ('attn', 'proj'), ('ffn', 'w1'), ('ffn', 'w2'), ('ffn', 'w3')] def add_lora(model, rank, alpha): """Freeze everything, then wrap every projection in the blocks with a LoRA detour.""" for p in model.parameters(): p.requires_grad_(False) for block in model.blocks: for part, name in TARGETS: parent = getattr(block, part) setattr(parent, name, LoRALinear(getattr(parent, name), rank, alpha)) return [p for p in model.parameters() if p.requires_grad] def merge_lora(model): for block in model.blocks: for part, name in TARGETS: parent = getattr(block, part) layer = getattr(parent, name) if isinstance(layer, LoRALinear): setattr(parent, name, layer.merge()) def main(): ap = argparse.ArgumentParser() ap.add_argument('--base', required=True) ap.add_argument('--chats', required=True) ap.add_argument('--data', default='data') ap.add_argument('--out', required=True) ap.add_argument('--rank', type=int, default=8) ap.add_argument('--alpha', type=float, default=16) ap.add_argument('--epochs', type=float, default=9) ap.add_argument('--batch', type=int, default=16) ap.add_argument('--lr', type=float, default=3e-3) ap.add_argument('--seed', type=int, default=0) args = ap.parse_args() torch.manual_seed(args.seed) device = get_device() os.makedirs(args.out, exist_ok=True) tok = Tokenizer.load(f'{args.data}/tokenizer.json') ck = torch.load(args.base, map_location='cpu') cfg = Config(**ck['config']) model = GPT(cfg) model.load_state_dict(ck['model']) trainable = add_lora(model, args.rank, args.alpha) model.to(device) n_train = sum(p.numel() for p in trainable) print(f'LoRA rank {args.rank}: {n_train / 1e3:.0f}k trainable of {model.num_params() / 1e6:.2f}M') chats = [json.loads(line) for line in open(args.chats)] random.Random(args.seed).shuffle(chats) n_val = max(8, len(chats) // 20) val_x, val_m = pack([render(tok, c['messages']) for c in chats[:n_val]], cfg.context) train_x, train_m = pack([render(tok, c['messages']) for c in chats[n_val:]], cfg.context) opt = torch.optim.AdamW(trainable, lr=args.lr, weight_decay=0.0) ctx = torch.autocast(device_type=device, dtype=torch.bfloat16) steps = max(1, int(args.epochs * len(train_x) / args.batch)) log = open(f'{args.out}/log.jsonl', 'w') log.write(json.dumps({'config': cfg.to_dict(), 'args': vars(args), 'params': model.num_params() - n_train, 'trainable': n_train, 'steps': steps, 'conversations': len(chats)}) + '\n') @torch.no_grad() def evaluate(): model.eval() x, m = val_x.to(device), val_m.to(device) with ctx: _, loss = model(x[:, :-1], x[:, 1:], loss_mask=m[:, 1:]) model.train() return loss.item() t0 = time.time() for step in range(steps + 1): if step % 10 == 0 or step == steps: rec = {'step': step, 'val': round(evaluate(), 4), 'time': round(time.time() - t0, 1)} print(json.dumps(rec), flush=True) log.write(json.dumps(rec) + '\n') if step == steps: break idx = torch.randint(0, len(train_x), (args.batch,)) x, m = train_x[idx].to(device), train_m[idx].to(device) with ctx: _, loss = model(x[:, :-1], x[:, 1:], loss_mask=m[:, 1:]) loss.backward() for g in opt.param_groups: g['lr'] = args.lr * lr_factor(step, steps, warmup=10, cooldown=0.5) opt.step() opt.zero_grad(set_to_none=True) log.write(json.dumps({'step': step, 'loss': round(loss.item(), 4)}) + '\n') log.close() adapter = {n: p.detach().cpu() for n, p in model.named_parameters() if p.requires_grad} torch.save({'rank': args.rank, 'alpha': args.alpha, 'adapter': adapter}, f'{args.out}/adapter.pt') merge_lora(model) torch.save({'config': cfg.to_dict(), 'model': model.state_dict(), 'step': steps}, f'{args.out}/ckpt_final.pt') if __name__ == '__main__': main()