Floquet Monodromy Optimizer / floquet_optimizer.py

Failed on benchmark

Raw ⬇ ZIP
 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())