Performs the SGD parameter update and stores :math:`v` in the optimizer state.
(self, gradient: mx.array, parameter: mx.array, state: dict)
| 270 | state["v"] = mx.zeros_like(parameter) |
| 271 | |
| 272 | def apply_single(self, gradient: mx.array, parameter: mx.array, state: dict): |
| 273 | """Performs the SGD parameter update and stores :math:`v` in the |
| 274 | optimizer state.""" |
| 275 | |
| 276 | if self.weight_decay != 0: |
| 277 | gradient += self.weight_decay * parameter |
| 278 | |
| 279 | if self.momentum <= 0: |
| 280 | return parameter - self.learning_rate.astype(gradient.dtype) * gradient |
| 281 | |
| 282 | v = self.momentum * state.get("v") |
| 283 | if self.dampening > 0: |
| 284 | v += (1 - self.dampening) * gradient |
| 285 | else: |
| 286 | v += gradient |
| 287 | |
| 288 | if self.nesterov: |
| 289 | update = gradient + self.momentum * v |
| 290 | else: |
| 291 | update = v |
| 292 | |
| 293 | state["v"] = v |
| 294 | return parameter - self.learning_rate.astype(gradient.dtype) * update |
| 295 | |
| 296 | |
| 297 | class RMSprop(Optimizer): |