Skip to content

Commit 552dbae

Browse files
committed
refactor: graft
1 parent ffcbcd3 commit 552dbae

File tree

1 file changed

+2
-2
lines changed

1 file changed

+2
-2
lines changed

pytorch_optimizer/optimizer/shampoo_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ class SGDGraft(Graft):
4747

4848
def __init__(self, var: torch.Tensor):
4949
super().__init__(var)
50-
self.momentum: torch.Tensor = torch.zeros_like(var, device=var.device)
50+
self.momentum: torch.Tensor = torch.zeros_like(var)
5151

5252
def update_momentum(self, update: torch.Tensor, beta1: float) -> torch.Tensor:
5353
r"""Update momentum."""
@@ -105,7 +105,7 @@ def add_statistics(self, grad: torch.Tensor, beta2: float) -> None:
105105

106106
def precondition_gradient(self, grad: torch.Tensor) -> torch.Tensor:
107107
r"""Get preconditioned gradient."""
108-
return grad / self.statistics.sqrt().add_(self.diagonal_eps)
108+
return grad.div(self.statistics.sqrt().add_(self.diagonal_eps))
109109

110110

111111
class BlockPartitioner:

0 commit comments

Comments
 (0)