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