Performs a single optimization step on the whole model or individual parameter. Args: closure: A closure which reevaluates the model and returns the loss. Incompatible with performing an optimization step on a single `param`. param: An individual para
(self, closure: Callable | None = None, param: Tensor | None = None)
| 137 | |
| 138 | @torch.no_grad() |
| 139 | def step(self, closure: Callable | None = None, param: Tensor | None = None): |
| 140 | """Performs a single optimization step on the whole model or individual parameter. |
| 141 | |
| 142 | Args: |
| 143 | closure: A closure which reevaluates the model and returns the loss. Incompatible with |
| 144 | performing an optimization step on a single `param`. |
| 145 | param: An individual parameter to perform a fused optimization step during the backward |
| 146 | pass. Requires optimizer to be initialized with `gradient_release=True` and model |
| 147 | hooks created with `register_gradient_release`. Incompatible with `closure`. |
| 148 | """ |
| 149 | loss = None |
| 150 | if closure is not None and param is None: |
| 151 | with torch.enable_grad(): |
| 152 | loss = closure() |
| 153 | |
| 154 | for group in self.param_groups: |
| 155 | params, grads, exp_avgs, exp_avg_sqs, eps_sqs, kahan_comps = [], [], [], [], [], [] |
| 156 | self._init_group(group, params, grads, exp_avgs, exp_avg_sqs, eps_sqs, kahan_comps) |
| 157 | |
| 158 | l1_norm, l2_norm = stableadamw( |
| 159 | params=params, |
| 160 | grads=grads, |
| 161 | exp_avgs=exp_avgs, |
| 162 | exp_avg_sqs=exp_avg_sqs, |
| 163 | eps_sqs=eps_sqs, |
| 164 | kahan_comps=kahan_comps, |
| 165 | lr=group["lr"], |
| 166 | beta1=group["beta1"], |
| 167 | beta2=group["beta2"], |
| 168 | weight_decay=group["weight_decay"], |
| 169 | eps=group["eps"], |
| 170 | step=group["step"], |
| 171 | decouple_lr=group["decouple_lr"], |
| 172 | max_lr=group["max_lr"], |
| 173 | kahan_sum=group["kahan_sum"], |
| 174 | return_norms=self.return_norms, |
| 175 | ) |
| 176 | |
| 177 | self.grad_norms["l1_norm"] = l1_norm |
| 178 | self.grad_norms["l2_norm"] = l2_norm |
| 179 | |
| 180 | return loss |
| 181 | |
| 182 | |
| 183 | def stableadamw( |
no test coverage detected