Floquet Monodromy Optimizer / floquet_optimizer.py
Failed on benchmark
1import torch
2
3
4class FloquetSGD(torch.optim.Optimizer):
5 """SGD with a repeating two-phase learning-rate schedule."""
6
7 def __init__(self, params, lr1, lr2, phase1_steps=1, phase2_steps=1,
8 momentum=0.0, weight_decay=0.0):
9 if lr1 < 0 or lr2 < 0:
10 raise ValueError("learning rates must be nonnegative")
11 if phase1_steps < 1 or phase2_steps < 1:
12 raise ValueError("phase lengths must be positive")
13 defaults = dict(lr1=float(lr1), lr2=float(lr2),
14 phase1_steps=int(phase1_steps),
15 phase2_steps=int(phase2_steps),
16 momentum=float(momentum),
17 weight_decay=float(weight_decay))
18 super().__init__(params, defaults)
19 self.step_count = 0
20
21 @property
22 def period(self):
23 g = self.param_groups[0]
24 return g["phase1_steps"] + g["phase2_steps"]
25
26 @property
27 def phase(self):
28 g = self.param_groups[0]
29 return 1 if self.step_count % self.period < g["phase1_steps"] else 2
30
31 @torch.no_grad()
32 def step(self, closure=None):
33 loss = None
34 if closure is not None:
35 with torch.enable_grad():
36 loss = closure()
37 for group in self.param_groups:
38 phase_offset = self.step_count % (group["phase1_steps"] + group["phase2_steps"])
39 lr = group["lr1"] if phase_offset < group["phase1_steps"] else group["lr2"]
40 momentum = group["momentum"]
41 decay = group["weight_decay"]
42 for p in group["params"]:
43 if p.grad is None:
44 continue
45 grad = p.grad
46 if decay:
47 grad = grad.add(p, alpha=decay)
48 if momentum:
49 state = self.state[p]
50 if "momentum_buffer" not in state:
51 buf = state["momentum_buffer"] = grad.detach().clone()
52 else:
53 buf = state["momentum_buffer"]
54 buf.mul_(momentum).add_(grad)
55 grad = buf
56 p.add_(grad, alpha=-lr)
57 self.step_count += 1
58 return loss
59
60
61def smoke_test():
62 torch.manual_seed(1053)
63 x = torch.tensor([[1.0], [2.0], [3.0]])
64 y = 2.0 * x
65 w = torch.nn.Parameter(torch.tensor([[0.0]]))
66 opt = FloquetSGD([w], lr1=0.1, lr2=0.02, phase1_steps=2, phase2_steps=1)
67 phases = []
68 for _ in range(6):
69 opt.zero_grad()
70 loss = ((x @ w - y) ** 2).mean()
71 loss.backward()
72 phases.append(opt.phase)
73 opt.step()
74 return {"phases_before_updates": phases, "final_weight": float(w),
75 "final_loss": float(((x @ w - y) ** 2).mean())}
76
77
78if __name__ == "__main__":
79 print(smoke_test())