MCPcopy Create free account
hub / github.com/ml-explore/mlx / test_adafactor

Method test_adafactor

python/tests/test_optimizers.py:269–287  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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 = {

Callers

nothing calls this directly

Calls 1

apply_gradientsMethod · 0.45

Tested by

no test coverage detected