import torch class FloquetSGD(torch.optim.Optimizer): """SGD with a repeating two-phase learning-rate schedule.""" def __init__(self, params, lr1, lr2, phase1_steps=1, phase2_steps=1, momentum=0.0, weight_decay=0.0): if lr1 < 0 or lr2 < 0: raise ValueError("learning rates must be nonnegative") if phase1_steps < 1 or phase2_steps < 1: raise ValueError("phase lengths must be positive") defaults = dict(lr1=float(lr1), lr2=float(lr2), phase1_steps=int(phase1_steps), phase2_steps=int(phase2_steps), momentum=float(momentum), weight_decay=float(weight_decay)) super().__init__(params, defaults) self.step_count = 0 @property def period(self): g = self.param_groups[0] return g["phase1_steps"] + g["phase2_steps"] @property def phase(self): g = self.param_groups[0] return 1 if self.step_count % self.period < g["phase1_steps"] else 2 @torch.no_grad() def step(self, closure=None): loss = None if closure is not None: with torch.enable_grad(): loss = closure() for group in self.param_groups: phase_offset = self.step_count % (group["phase1_steps"] + group["phase2_steps"]) lr = group["lr1"] if phase_offset < group["phase1_steps"] else group["lr2"] momentum = group["momentum"] decay = group["weight_decay"] for p in group["params"]: if p.grad is None: continue grad = p.grad if decay: grad = grad.add(p, alpha=decay) if momentum: state = self.state[p] if "momentum_buffer" not in state: buf = state["momentum_buffer"] = grad.detach().clone() else: buf = state["momentum_buffer"] buf.mul_(momentum).add_(grad) grad = buf p.add_(grad, alpha=-lr) self.step_count += 1 return loss def smoke_test(): torch.manual_seed(1053) x = torch.tensor([[1.0], [2.0], [3.0]]) y = 2.0 * x w = torch.nn.Parameter(torch.tensor([[0.0]])) opt = FloquetSGD([w], lr1=0.1, lr2=0.02, phase1_steps=2, phase2_steps=1) phases = [] for _ in range(6): opt.zero_grad() loss = ((x @ w - y) ** 2).mean() loss.backward() phases.append(opt.phase) opt.step() return {"phases_before_updates": phases, "final_weight": float(w), "final_loss": float(((x @ w - y) ** 2).mean())} if __name__ == "__main__": print(smoke_test())