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