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

Method test_tree

python/tests/test_vmap.py:111–165  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 1

vmapMethod · 0.45

Tested by

no test coverage detected