File size: 3,220 Bytes
ea3a71e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 | """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))
|