| """Federated / decentralized DPO on the paper's actual setup: DistilGPT-2 (82M) |
| trained with the true DPO loss on Stanford Human Preferences pairs, split across |
| N=5 clients. This replaces the 64-dim log-linear proxy used previously. |
| """ |
| import os, json, time, math, copy, argparse |
| import numpy as np, torch, torch.nn.functional as F |
| from transformers import AutoTokenizer, AutoModelForCausalLM |
|
|
| DEV = "mps" if torch.backends.mps.is_available() else "cpu" |
| BETA = 0.1 |
| MODEL = "distilgpt2" |
|
|
|
|
| def load_pairs(n_pairs, maxlen=128, seed=0): |
| from datasets import load_dataset |
| ds = load_dataset("stanfordnlp/SHP", split="train", streaming=True) |
| tok = AutoTokenizer.from_pretrained(MODEL) |
| tok.pad_token = tok.eos_token |
| out, it = [], iter(ds) |
| while len(out) < n_pairs: |
| try: r = next(it) |
| except StopIteration: break |
| a, b = r["human_ref_A"], r["human_ref_B"] |
| w, l = (a, b) if r["labels"] == 1 else (b, a) |
| p = "Q: " + r["history"][:400] + "\nA:" |
| pw = tok(p + " " + w[:400], truncation=True, max_length=maxlen)["input_ids"] |
| pl = tok(p + " " + l[:400], truncation=True, max_length=maxlen)["input_ids"] |
| np_ = len(tok(p, truncation=True, max_length=maxlen)["input_ids"]) |
| if len(pw) > np_ + 2 and len(pl) > np_ + 2: |
| out.append({"w": pw, "l": pl, "np": np_, "domain": r["domain"]}) |
| return out, tok |
|
|
|
|
| def batch_logps(model, seqs, nps, pad): |
| L = max(len(s) for s in seqs) |
| ids = torch.full((len(seqs), L), pad, dtype=torch.long) |
| msk = torch.zeros((len(seqs), L), dtype=torch.bool) |
| for i, s in enumerate(seqs): |
| ids[i, :len(s)] = torch.tensor(s); msk[i, nps[i]:len(s)] = True |
| ids, msk = ids.to(DEV), msk.to(DEV) |
| logits = model(ids).logits[:, :-1] |
| tgt, m = ids[:, 1:], msk[:, 1:] |
| lp = torch.log_softmax(logits.float(), -1).gather(2, tgt.unsqueeze(-1)).squeeze(-1) |
| return (lp * m).sum(-1) |
|
|
|
|
| def dpo_loss(model, ref, batch, pad): |
| sw = [b["w"] for b in batch]; sl = [b["l"] for b in batch] |
| nw = [b["np"] for b in batch]; nl = [b["np"] for b in batch] |
| pw = batch_logps(model, sw, nw, pad); pl = batch_logps(model, sl, nl, pad) |
| with torch.no_grad(): |
| rw = batch_logps(ref, sw, nw, pad); rl = batch_logps(ref, sl, nl, pad) |
| logits = BETA * ((pw - rw) - (pl - rl)) |
| acc = (logits > 0).float().mean().item() |
| return -F.logsigmoid(logits).mean(), acc |
|
|
|
|
| if __name__ == "__main__": |
| t0 = time.time() |
| pairs, tok = load_pairs(40) |
| print("loaded %d pairs in %.1fs; domains=%s" % (len(pairs), time.time()-t0, |
| sorted({p['domain'] for p in pairs})[:5]), flush=True) |
| m = AutoModelForCausalLM.from_pretrained(MODEL).to(DEV) |
| ref = AutoModelForCausalLM.from_pretrained(MODEL).to(DEV).eval() |
| for p in ref.parameters(): p.requires_grad_(False) |
| opt = torch.optim.SGD(m.parameters(), lr=1e-4) |
| t0 = time.time() |
| for step in range(5): |
| loss, acc = dpo_loss(m, ref, pairs[step*4:(step+1)*4], tok.pad_token_id) |
| opt.zero_grad(); loss.backward(); opt.step() |
| print(" step %d loss=%.4f acc=%.2f (%.2fs)" % (step, loss.item(), acc, time.time()-t0), flush=True) |
| print("per-step: %.2fs on %s" % ((time.time()-t0)/5, DEV)) |
|
|