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