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

Method test_custom_function

python/tests/test_autograd.py:603–685  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

601 self.assertTrue(mx.array_equal(out, expected))
602
603 def test_custom_function(self):
604 # Make a custom function
605 my_exp = mx.custom_function(mx.exp)
606
607 # Ensure everything works
608 dy = mx.grad(my_exp)(mx.array(1.0))
609 self.assertTrue(mx.allclose(dy, mx.exp(mx.array(1.0))))
610 (ex,), (dex,) = mx.jvp(my_exp, [mx.array(1.0)], [mx.array(1.0)])
611 self.assertTrue(mx.allclose(dex, mx.exp(mx.array(1.0))))
612 self.assertTrue(mx.allclose(ex, dex))
613 ex = mx.vmap(my_exp)(mx.ones(10))
614 self.assertTrue(mx.allclose(ex, mx.exp(mx.ones(10))))
615
616 # Ensure that the vjp is being overriden but everything else still
617 # works.
618 @my_exp.vjp
619 def my_exp_vjp(x, dx, ex):
620 return mx.ones_like(x) * 42
621
622 dy = mx.grad(my_exp)(mx.array(1.0))
623 self.assertTrue(mx.allclose(dy, mx.array(42.0)))
624 (ex,), (dex,) = mx.jvp(my_exp, [mx.array(1.0)], [mx.array(1.0)])
625 self.assertTrue(mx.allclose(dex, mx.exp(mx.array(1.0))))
626 self.assertTrue(mx.allclose(ex, dex))
627 ex = mx.vmap(my_exp)(mx.ones(10))
628 self.assertTrue(mx.allclose(ex, mx.exp(mx.ones(10))))
629
630 # Ensure that setting the jvp and vmap also works.
631 @my_exp.jvp
632 def my_exp_jvp(x, dx):
633 return mx.ones_like(x) * 7 * dx
634
635 @my_exp.vmap
636 def my_exp_vmap(x, axis):
637 return mx.ones_like(x) * 3, axis
638
639 dy = mx.grad(my_exp)(mx.array(1.0))
640 self.assertTrue(mx.allclose(dy, mx.array(42.0)))
641 (ex,), (dex,) = mx.jvp(my_exp, [mx.array(1.0)], [mx.array(1.0)])
642 self.assertTrue(mx.allclose(dex, mx.array(7.0)))
643 self.assertTrue(mx.allclose(ex, mx.exp(mx.array(1.0))))
644 ex = mx.vmap(my_exp)(mx.ones(10))
645 self.assertTrue(mx.allclose(ex, 3 * mx.ones(10)))
646
647 # Test pytrees
648 @mx.custom_function
649 def my_double(params):
650 return {"out": 2 * params["x"] * params["y"]}
651
652 dy = mx.grad(lambda p: my_double(p)["out"].sum())(
653 {"x": mx.ones(2), "y": mx.ones(2)}
654 )
655 self.assertTrue(mx.allclose(dy["x"], mx.ones(2) * 2))
656 self.assertTrue(mx.allclose(dy["y"], mx.ones(2) * 2))
657
658 @my_double.vjp
659 def random_grads(primals, cotangents, outputs):
660 return {"x": mx.zeros_like(primals["x"]), "y": mx.ones_like(primals["y"])}

Callers

nothing calls this directly

Calls 3

arrayMethod · 0.60
jvpMethod · 0.45
vmapMethod · 0.45

Tested by

no test coverage detected