"""LoRA SFT for TinyLiquid, per the small-model adaptation recipe (LoRA paper). Freezes the pretrained base, trains low-rank adapters on the gated-MLP linears plus persona embeddings, with a KL anchor to the base and a TinyStories PPL guard. Best checkpoint is selected by masked SFT holdout loss while PPL < guard. Saved checkpoints are FOLDED back into standard model keys (no lora_* in the state dict), so hf/export_hf.py works unchanged. Usage: .venv/bin/python train/train_lora.py --base ckpt/nlp --data data/sft_mix_v5.jsonl \ --ckpt ckpt/v5_lora --epochs 2 --lr 3e-4 --r 16 --kl 0.05 """ import argparse, json, math, random, time from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from model.config import TinyLiquidConfig, CONFIGS from model.utils import latest_ckpt from model.tiny_liquid import TinyLiquid from data.tokenizer import load_tokenizer def resolve_ckpt(path): p = Path(path) if p.is_file(): return p ck = latest_ckpt(p) assert ck, f"no checkpoints in {path}" return ck USER_T, ASST_T, EOT_T = "<|user|>", "<|assistant|>", "<|endoftext|>" PERSONA_T = {"analyst": "<|analyst|>", "skeptic": "<|skeptic|>", "spock": "<|analyst|>", "none": ""} P_IDS = {"analyst": 1, "skeptic": 2, "spock": 1, "none": 0} class LoRALinear(nn.Module): def __init__(self, base: nn.Linear, r: int, alpha: float, dropout: float): super().__init__() self.base = base for p in base.parameters(): p.requires_grad = False out_f, in_f = base.weight.shape self.lora_a = nn.Parameter(torch.empty(in_f, r)) self.lora_b = nn.Parameter(torch.zeros(r, out_f)) nn.init.kaiming_uniform_(self.lora_a, a=math.sqrt(5)) self.dropout = nn.Dropout(dropout) self.scale = alpha / max(1, r) def forward(self, x): return self.base(x) + (self.dropout(x) @ self.lora_a @ self.lora_b) * self.scale def wrap_lora(model: TinyLiquid, r: int, alpha: float, dropout: float): wrapped = [] for name, mod in list(model.named_modules()): if isinstance(mod, nn.Linear) and not name.endswith("lm_head"): lora = LoRALinear(mod, r, alpha, dropout) parts = name.split(".") parent = model for p in parts[:-1]: parent = parent._modules[p] if isinstance(parent, nn.Module) else getattr(parent, p) parent._modules[parts[-1]] = lora wrapped.append((name, lora)) return wrapped def fold_state_dict(sd, wrapped): out = {} for k, v in sd.items(): if any(k.startswith(n + ".") and not k.startswith(n + ".base.") for n, _ in wrapped): continue # lora_a / lora_b matched = False for name, _ in wrapped: if k.startswith(name + ".base."): out[name + "." + k.split(".base.", 1)[1]] = v.clone() matched = True break if not matched: out[k] = v.clone() for name, lora in wrapped: delta = (lora.lora_a @ lora.lora_b).t() * lora.scale out[name + ".weight"] = out[name + ".weight"] + delta.detach() return out def parse_args(): ap = argparse.ArgumentParser() ap.add_argument("--base", default="ckpt/nlp") ap.add_argument("--resume", default="", help="resume from latest ckpt in this dir") ap.add_argument("--data", default="data/sft_mix_v5.jsonl") ap.add_argument("--tok", default="data/tokenizer.json") ap.add_argument("--ckpt", default="ckpt/v5_lora") ap.add_argument("--val-bin", default="data/valid.bin") ap.add_argument("--replay-bin", default="", help="tokenized bin to mix as fluency replay (raw full-loss items)") ap.add_argument("--replay-ratio", type=float, default=0.5, help="fraction of replay items in the train mixture (0..1)") ap.add_argument("--epochs", type=int, default=2) ap.add_argument("--batch", type=int, default=8) ap.add_argument("--seq", type=int, default=256) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--r", type=int, default=16) ap.add_argument("--alpha", type=float, default=32.0) ap.add_argument("--dropout", type=float, default=0.05) ap.add_argument("--kl", type=float, default=0.05) ap.add_argument("--eval-every", type=int, default=25) ap.add_argument("--log-every", type=int, default=25) ap.add_argument("--ppl-guard", type=float, default=60.0) ap.add_argument("--resume-best-sft", type=float, default=None) ap.add_argument("--resume-best-ppl", type=float, default=None) ap.add_argument("--val-batches", type=int, default=2) ap.add_argument("--seed", type=int, default=17) ap.add_argument("--threads", type=int, default=4) return ap.parse_args() def tokenize_example(tok, ex, seq, u_id, a_id, eot_id): if "raw" in ex: ids = tok.encode(ex["raw"]).ids + [eot_id] x = torch.tensor(ids[:-1], dtype=torch.long) y = torch.tensor(ids[1:], dtype=torch.long) return x[:seq], y[:seq], torch.ones_like(y[:seq], dtype=torch.bool), 0 persona_name = ex.get("persona", "analyst") persona = PERSONA_T.get(persona_name, PERSONA_T["analyst"]) p_id = P_IDS.get(persona_name, 1) p_ids = tok.encode(persona).ids if persona else [] ids = p_ids + [u_id] + tok.encode(ex["user"]).ids + [a_id] + tok.encode(ex["assistant"]).ids + [eot_id] if len(ids) > seq: return None asst_start = len(p_ids) + 1 + len(tok.encode(ex["user"]).ids) + 1 x = torch.tensor(ids[:-1], dtype=torch.long) y = torch.tensor(ids[1:], dtype=torch.long) mask = torch.zeros_like(y, dtype=torch.bool) mask[asst_start - 1:] = True if int(mask.sum()) < 16: return None return x, y, mask, p_id def collate(items, seq): xs, ys, ms, ps = [], [], [], [] for x, y, m, p in items: xs.append(F.pad(x, (0, seq - x.shape[0]), value=0)) ys.append(F.pad(y, (0, seq - y.shape[0]), value=0)) ms.append(F.pad(m, (0, seq - m.shape[0]), value=False)) ps.append(p) return torch.stack(xs), torch.stack(ys), torch.stack(ms), torch.tensor(ps, dtype=torch.long) @torch.no_grad() def val_ppl(model, val_bin, batch=4, seq=64, n_batches=2, seed=0): mm = np.memmap(val_bin, dtype=np.uint16, mode="r") total, cnt = 0.0, 0 rng = np.random.RandomState(seed) n = (len(mm) - 1) // seq for _ in range(n_batches): s = int(rng.randint(0, n - batch)) buf = torch.stack([torch.from_numpy(mm[s * seq + i * seq: s * seq + i * seq + seq].astype(np.int64)) for i in range(batch)]) x, y = buf[:, :-1], buf[:, 1:] loss = F.cross_entropy(model(x).reshape(-1, model.cfg.vocab_size), y.reshape(-1)) total += loss.item() * y.numel(); cnt += y.numel() return float(np.exp(total / cnt)) def main(): args = parse_args() torch.set_num_threads(args.threads) torch.manual_seed(args.seed); random.seed(args.seed) rng = random.Random(args.seed) tok = load_tokenizer(args.tok) u_id, a_id, eot_id = tok.token_to_id(USER_T), tok.token_to_id(ASST_T), tok.token_to_id(EOT_T) assert None not in (u_id, a_id, eot_id) raw = [json.loads(l) for l in open(args.data, encoding="utf-8") if l.strip()] teacher_path = resolve_ckpt(args.base) model_path = resolve_ckpt(args.resume) if args.resume else teacher_path resume_ck = torch.load(model_path, map_location="cpu") if args.resume else None base = torch.load(model_path, map_location="cpu") config = base.get("config") or CONFIGS["tiny10m"] cfg = TinyLiquidConfig(vocab_size=tok.get_vocab_size(), **{k: v for k, v in config.items() if k != "vocab_size"}) cfg.mtp_heads = 0 # MTP is pretrain-only; post-training has no MTP heads model = TinyLiquid(cfg); model.load_state_dict(base["model"], strict=False) teacher_sd = torch.load(teacher_path, map_location="cpu")["model"] teacher = None if args.kl > 0: teacher = TinyLiquid(cfg); teacher.load_state_dict(teacher_sd, strict=False); teacher.eval() for p in teacher.parameters(): p.requires_grad = False wrapped = wrap_lora(model, args.r, args.alpha, args.dropout) for p in model.parameters(): p.requires_grad = False for p in model.persona_emb.parameters(): p.requires_grad = True for _, lora in wrapped: lora.base.weight.requires_grad = False lora.lora_a.requires_grad = True lora.lora_b.requires_grad = True trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=args.lr, betas=(0.9, 0.95), weight_decay=0.02) print(f"base {model_path.name} | lora adapters {len(wrapped)} | trainable {trainable:,}", flush=True) items_all = [tokenize_example(tok, e, args.seq, u_id, a_id, eot_id) for e in raw] items_all = [i for i in items_all if i is not None] rng.shuffle(items_all) n_eval = min(128, max(16, len(items_all) // 12)) eval_items, train_items = items_all[:n_eval], items_all[n_eval:] if args.replay_bin: mm = np.memmap(args.replay_bin, dtype=np.uint16, mode="r") n = (len(mm) - 1) // args.seq gold_n = max(1, len(train_items)) replay_n = int(gold_n * args.replay_ratio / max(1e-9, 1.0 - args.replay_ratio)) rr = np.random.RandomState(args.seed + 1) for _ in range(replay_n): s = int(rr.randint(0, n)) w = torch.from_numpy(mm[s * args.seq: (s + 1) * args.seq].astype(np.int64)) x, y = w[:-1], w[1:] train_items.append((x, y, torch.ones_like(y, dtype=torch.bool), 0)) print(f"replay: {replay_n} raw items from {args.replay_bin} (mixture ratio {args.replay_ratio:.2f})", flush=True) print(f"train {len(train_items)} eval {len(eval_items)} filtered {len(raw) - len(items_all)}", flush=True) out = Path(args.ckpt); out.mkdir(parents=True, exist_ok=True) best_score = args.resume_best_sft if args.resume_best_sft is not None else float("inf") best_ppl = args.resume_best_ppl if args.resume_best_ppl is not None else float("inf") step = (resume_ck or {}).get("step", 0) start_iter = (resume_ck or {}).get("iter", step) t0 = time.time() total_steps = (len(train_items) // args.batch) * args.epochs if step: print(f"resuming from {model_path} at step {step}/{total_steps} iter {start_iter}", flush=True) def run_eval(): model.eval() total, n = 0.0, 0 for i in range(0, len(eval_items), args.batch): x, y, m, p = collate(eval_items[i:i + args.batch], args.seq) with torch.no_grad(): logits = model(x, persona_ids=p).reshape(-1, model.cfg.vocab_size) loss = F.cross_entropy(logits, y.reshape(-1), reduction="none") loss = (loss * m.reshape(-1)).sum() / m.sum() total += loss.item() * m.sum().item(); n += m.sum().item() ppl = val_ppl(model, args.val_bin, n_batches=args.val_batches, seed=args.seed + step) model.train() return total / n, ppl def save(path, tag=""): sd = fold_state_dict(model.state_dict(), wrapped) torch.save({"model": sd, "step": step, "iter": iter_no, "config": cfg.__dict__, "tag": tag}, str(path)) model.train() iter_no = 0 for ep in range(args.epochs): rng.shuffle(train_items) usable = len(train_items) - len(train_items) % args.batch for i in range(0, usable, args.batch): iter_no += 1 if iter_no <= start_iter: continue step += 1 x, y, m, p = collate(train_items[i:i + args.batch], args.seq) opt.zero_grad(set_to_none=True) logits = model(x, persona_ids=p) sft_loss = F.cross_entropy(logits.reshape(-1, model.cfg.vocab_size), y.reshape(-1), reduction="none") sft_loss = (sft_loss * m.reshape(-1)).sum() / m.sum() loss = sft_loss if teacher is not None: with torch.no_grad(): t_logits = teacher(x, persona_ids=p) kl = F.kl_div(F.log_softmax(logits.float(), dim=-1), F.softmax(t_logits.float(), dim=-1), reduction="none").sum(dim=-1) loss = loss + args.kl * (kl * m).sum() / m.sum() loss.backward() torch.nn.utils.clip_grad_norm_([p for p in model.parameters() if p.requires_grad], 0.5) opt.step() if step % args.log_every == 0: print(f"step {step}/{total_steps} loss {loss.item():.4f} sft {sft_loss.item():.4f} " f"{args.batch * args.seq * args.log_every / max(1e-6, time.time() - t0):.0f} tok/s", flush=True) t0 = time.time() if step % args.eval_every == 0: sft_vl, ppl = run_eval() tag = "" if ppl < best_ppl: best_ppl = ppl; save(out / "best_ppl.pt", tag="best_ppl"); tag += " [best ppl]" if ppl <= args.ppl_guard and sft_vl < best_score: best_score = sft_vl; save(out / "best.pt", tag="best"); tag += " [new best]" save(out / f"model_{step}.pt", tag=f"step{step}") try: sp = tok.encode("<|analyst|><|user|>Find discrepancies between: Account A: The meeting ended at 11am. Account B: The meeting ended at noon.<|assistant|>").ids with torch.no_grad(): sout = tok.decode(model.generate(tok, sp, persona_id=1, max_new=50, temperature=0.35, top_k=20, repetition_penalty=1.25, no_repeat_ngram_size=4)[len(sp):]).replace("\n", " ").strip()[:180] print(f" [eval {step}] sft_val_loss {sft_vl:.4f} val_ppl {ppl:.2f}{tag}", flush=True) print(f" sample: {sout}", flush=True) except Exception: print(f" [eval {step}] sft_val_loss {sft_vl:.4f} val_ppl {ppl:.2f}{tag}", flush=True) sd = fold_state_dict(model.state_dict(), wrapped) torch.save({"model": sd, "step": step, "config": cfg.__dict__, "tag": "final"}, out / "model_final.pt") print(f"done -> {out} best_sft={best_score:.4f} best_ppl={best_ppl:.2f}", flush=True) if __name__ == "__main__": main()