| 109 | self.assertTrue(mx.array_equal(out, op(x, y.T).T)) |
| 110 | |
| 111 | def test_tree(self): |
| 112 | def my_fun(tree): |
| 113 | return (tree["a"] + tree["b"][0]) * tree["b"][1] |
| 114 | |
| 115 | tree = { |
| 116 | "a": mx.random.uniform(shape=(2, 4)), |
| 117 | "b": ( |
| 118 | mx.random.uniform(shape=(2, 4)), |
| 119 | mx.random.uniform(shape=(2, 4)), |
| 120 | ), |
| 121 | } |
| 122 | out = mx.vmap(my_fun)(tree) |
| 123 | expected = my_fun(tree) |
| 124 | self.assertTrue(mx.array_equal(out, my_fun(tree))) |
| 125 | |
| 126 | with self.assertRaises(ValueError): |
| 127 | mx.vmap(my_fun, in_axes={"a": 0, "b": ((0, 0), 0)}, out_axes=0)(tree) |
| 128 | |
| 129 | out = mx.vmap(my_fun, in_axes={"a": 0, "b": 0}, out_axes=0)(tree) |
| 130 | self.assertTrue(mx.array_equal(out, my_fun(tree))) |
| 131 | |
| 132 | out = mx.vmap(my_fun, in_axes={"a": 0, "b": (0, 0)}, out_axes=0)(tree) |
| 133 | self.assertTrue(mx.array_equal(out, my_fun(tree))) |
| 134 | |
| 135 | tree = { |
| 136 | "a": mx.random.uniform(shape=(2, 4)), |
| 137 | "b": ( |
| 138 | mx.random.uniform(shape=(4, 2)), |
| 139 | mx.random.uniform(shape=(4, 2)), |
| 140 | ), |
| 141 | } |
| 142 | out = mx.vmap(my_fun, in_axes={"a": 0, "b": (1, 1)}, out_axes=0)(tree) |
| 143 | expected = (tree["a"] + tree["b"][0].T) * tree["b"][1].T |
| 144 | self.assertTrue(mx.array_equal(out, expected)) |
| 145 | |
| 146 | def my_fun(x, y): |
| 147 | return {"a": x + y, "b": x * y} |
| 148 | |
| 149 | x = mx.random.uniform(shape=(2, 4)) |
| 150 | y = mx.random.uniform(shape=(2, 4)) |
| 151 | out = mx.vmap(my_fun, in_axes=0, out_axes=0)(x, y) |
| 152 | expected = my_fun(x, y) |
| 153 | self.assertTrue(mx.array_equal(out["a"], expected["a"])) |
| 154 | self.assertTrue(mx.array_equal(out["b"], expected["b"])) |
| 155 | |
| 156 | with self.assertRaises(ValueError): |
| 157 | mx.vmap(my_fun, in_axes=0, out_axes=(0, 1))(x, y) |
| 158 | |
| 159 | with self.assertRaises(ValueError): |
| 160 | mx.vmap(my_fun, in_axes=0, out_axes={"a": 0, "c": 1})(x, y) |
| 161 | |
| 162 | out = mx.vmap(my_fun, in_axes=0, out_axes={"a": 1, "b": 0})(x, y) |
| 163 | expected = my_fun(x, y) |
| 164 | self.assertTrue(mx.array_equal(out["a"].T, expected["a"])) |
| 165 | self.assertTrue(mx.array_equal(out["b"], expected["b"])) |
| 166 | |
| 167 | def test_vmap_indexing(self): |
| 168 | x = mx.arange(16).reshape(2, 2, 2, 2) |