Download code/claim1_real_scope.py from SabaPivot/repro-distributed-direct-preference-optimization: direct link, hf CLI and curl.
- Browser
- Download file 2.16 kB
-
https://huggingface.co/spaces/SabaPivot/repro-distributed-direct-preference-optimization/resolve/main/code/claim1_real_scope.py
- Command line
-
hf download hf://spaces/SabaPivot/repro-distributed-direct-preference-optimization/code/claim1_real_scope.py
-
curl -L -o claim1_real_scope.py https://huggingface.co/spaces/SabaPivot/repro-distributed-direct-preference-optimization/resolve/main/code/claim1_real_scope.py
2.16 kB
| """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() | |