@@ -66,19 +66,15 @@ def __setstate__(self, state):
6666 if not p_state :
6767 continue
6868 if 'step' in p_state :
69- p_state ['step' ] = _init_scalar (float (p_state ['step' ]), device = 'cpu' )
70- if (
71- 'exp_avg_lr_1' in p_state
72- and torch .is_tensor (group ['lr' ])
73- and not torch .is_tensor (p_state ['exp_avg_lr_1' ])
74- ):
75- p_state ['exp_avg_lr_1' ] = torch .tensor (
76- float (p_state ['exp_avg_lr_1' ]),
69+ p_state ['step' ] = _init_scalar (p_state ['step' ], device = 'cpu' )
70+ if 'exp_avg_lr_1' in p_state and torch .is_tensor (group ['lr' ]):
71+ p_state ['exp_avg_lr_1' ] = _init_scalar (
72+ p_state ['exp_avg_lr_1' ],
7773 dtype = group ['lr' ].dtype ,
7874 device = group ['lr' ].device ,
7975 )
8076 if 'exp_avg_lr_2' in p_state :
81- p_state ['exp_avg_lr_2' ] = _init_scalar (float ( p_state ['exp_avg_lr_2' ]) , device = 'cpu' )
77+ p_state ['exp_avg_lr_2' ] = _init_scalar (p_state ['exp_avg_lr_2' ], device = 'cpu' )
8278
8379 @torch .no_grad ()
8480 def step (self , closure = None ):
@@ -129,9 +125,10 @@ def step(self, closure=None):
129125
130126 # 1 - beta1 ** state['step']
131127 if torch .is_tensor (group ['lr' ]):
128+ lr_safe = torch .where (group ['lr' ] != 0. , group ['lr' ], torch .ones_like (group ['lr' ]))
132129 bias_correction1 = torch .where (
133130 group ['lr' ] != 0. ,
134- state ['exp_avg_lr_1' ] / group [ 'lr' ] ,
131+ state ['exp_avg_lr_1' ] / lr_safe ,
135132 torch .ones_like (group ['lr' ]),
136133 )
137134 else :
0 commit comments