Performs a single optimization step. Args: closure (Callable, optional): A closure that reevaluates the model and returns the loss.
(self, closure=None)
| 54 | |
| 55 | @_use_grad_for_differentiable |
| 56 | def step(self, closure=None): |
| 57 | """Performs a single optimization step. |
| 58 | |
| 59 | Args: |
| 60 | closure (Callable, optional): A closure that reevaluates the model |
| 61 | and returns the loss. |
| 62 | """ |
| 63 | loss = None |
| 64 | if closure is not None: |
| 65 | with torch.enable_grad(): |
| 66 | loss = closure() |
| 67 | |
| 68 | for group in self.param_groups: |
| 69 | params_with_grad = [] |
| 70 | d_p_list = [] |
| 71 | momentum_buffer_list = [] |
| 72 | |
| 73 | has_sparse_grad = self._init_group(group, params_with_grad, d_p_list, momentum_buffer_list) |
| 74 | |
| 75 | sgd(params_with_grad, |
| 76 | d_p_list, |
| 77 | momentum_buffer_list, |
| 78 | weight_decay=group['weight_decay'], |
| 79 | momentum=group['momentum'], |
| 80 | lr=group['lr'], |
| 81 | dampening=group['dampening'], |
| 82 | nesterov=group['nesterov'], |
| 83 | maximize=group['maximize'], |
| 84 | has_sparse_grad=has_sparse_grad, |
| 85 | foreach=group['foreach']) |
| 86 | |
| 87 | # update momentum_buffers in state |
| 88 | for p, momentum_buffer in zip(params_with_grad, momentum_buffer_list): |
| 89 | state = self.state[p] |
| 90 | state['momentum_buffer'] = momentum_buffer |
| 91 | |
| 92 | return loss |
| 93 | |
| 94 | |
| 95 | SGD.__doc__ = r"""Implements stochastic gradient descent (optionally with momentum). |