File size: 2,163 Bytes
dd90a4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Focused real-model FedDPO scope run for registered claim 1."""
import json
import time

import torch
from transformers import AutoModelForCausalLM

from dpo_real import dpo_loss, DEV, MODEL
from fed_real import build_clients, fed_run, evaluate


OUT = "outputs/claim1_real_scope.json"
R = 10
S = 3
LR = 2e-5


def gradient_observation(model, reference, clients, tok):
    model.zero_grad(set_to_none=True)
    losses = []
    for client in clients:
        loss, _ = dpo_loss(model, reference, client[:4], tok.pad_token_id)
        losses.append(loss)
    pooled = torch.stack(losses).mean()
    pooled.backward()
    norm_sq = 0.0
    for parameter in model.parameters():
        if parameter.grad is not None:
            norm_sq += float((parameter.grad.detach().float() ** 2).sum().item())
    model.zero_grad(set_to_none=True)
    return norm_sq, float(pooled.detach().item())


def main():
    started = time.time()
    clients, names, tok = build_clients()
    base = AutoModelForCausalLM.from_pretrained(MODEL)
    reference = AutoModelForCausalLM.from_pretrained(MODEL).to(DEV).eval()
    for parameter in reference.parameters():
        parameter.requires_grad_(False)
    rows = []
    for E in (1, 3, 6):
        model, _ = fed_run(base, reference, clients, tok, S=S, R=R, E=E, lr=LR, seed=0)
        loss, accuracy = evaluate(model, reference, clients, tok.pad_token_id, nb=3)
        grad2, pooled_loss = gradient_observation(model.to(DEV), reference, clients, tok)
        rows.append({"E": E, "S": S, "R": R, "lr": LR,
                     "final_dpo_loss": float(loss), "accuracy": float(accuracy),
                     "pooled_gradient_norm_sq": grad2, "pooled_dpo_loss": pooled_loss})
        print(json.dumps(rows[-1]), flush=True)
    payload = {"model": "distilgpt2 (82M)", "dataset": "stanfordnlp/SHP",
               "clients": dict(zip(names, [len(c) for c in clients])),
               "algorithm": "FedDPO with client sampling S=3 and R=10",
               "rows": rows, "elapsed_seconds": time.time() - started}
    with open(OUT, "w") as handle:
        json.dump(payload, handle, indent=2)


if __name__ == "__main__":
    main()