| 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"])} |