(self)
| 267 | ) |
| 268 | |
| 269 | def test_adafactor(self): |
| 270 | x = mx.zeros((5, 5)) |
| 271 | params = {"x": x} |
| 272 | grad = {"x": mx.ones_like(x)} |
| 273 | optimizer = opt.Adafactor() |
| 274 | for _ in range(2): |
| 275 | xp = optimizer.apply_gradients(grad, params) |
| 276 | self.assertEqual(xp["x"].dtype, x.dtype) |
| 277 | self.assertEqual(xp["x"].shape, x.shape) |
| 278 | |
| 279 | x = mx.zeros((5, 5), mx.float16) |
| 280 | params = {"x": x} |
| 281 | grad = {"x": mx.ones_like(x)} |
| 282 | optimizer = opt.Adafactor() |
| 283 | for _ in range(2): |
| 284 | xp = optimizer.apply_gradients(grad, params) |
| 285 | self.assertEqual(xp["x"].dtype, x.dtype) |
| 286 | self.assertEqual(xp["x"].shape, x.shape) |
| 287 | self.assertEqual(optimizer.state["step"], 2) |
| 288 | |
| 289 | def test_muon(self): |
| 290 | params = { |
nothing calls this directly
no test coverage detected