Skip to content

Commit ecc4530

Browse files
committed
fix: Kate optimizer
1 parent 4126a53 commit ecc4530

File tree

1 file changed

+5
-4
lines changed

1 file changed

+5
-4
lines changed

pytorch_optimizer/optimizer/kate.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -100,15 +100,16 @@ def step(self, closure: CLOSURE = None) -> LOSS:
100100
grad_p2 = grad * grad
101101

102102
m, b = state['m'], state['b']
103-
b.add_(grad_p2)
103+
b.mul_(b).add_(grad_p2)
104104

105105
de_nom = b.add(group['eps'])
106106

107-
m.add_(grad, alpha=group['eta']).add_(grad / de_nom)
107+
m.mul_(m).add_(grad_p2, alpha=group['eta']).add_(grad / de_nom).sqrt_()
108108

109-
update = m.sqrt()
110-
update.mul_(grad).div_(de_nom)
109+
update = m.mul(grad).div_(b)
111110

112111
p.add_(update, alpha=-group['lr'])
113112

113+
b.sqrt_()
114+
114115
return loss

0 commit comments

Comments
 (0)