| 47 | self.assertTrue(mx.array_equal(v1, v2)) |
| 48 | |
| 49 | def test_supported_trees(self): |
| 50 | |
| 51 | from typing import NamedTuple |
| 52 | |
| 53 | class Vector(tuple): |
| 54 | pass |
| 55 | |
| 56 | class Params(NamedTuple): |
| 57 | m: mx.array |
| 58 | b: mx.array |
| 59 | |
| 60 | list1 = [mx.array([0, 1]), mx.array(2)] |
| 61 | tuple1 = (mx.array([0, 1]), mx.array(2)) |
| 62 | vector1 = Vector([mx.array([0, 1]), mx.array(2)]) |
| 63 | params1 = Params(m=mx.array([0, 1]), b=mx.array(2)) |
| 64 | dict1 = {"m": mx.array([0, 1]), "b": mx.array(2)} |
| 65 | |
| 66 | add_one = lambda x: x + 1 |
| 67 | |
| 68 | list2 = mlx.utils.tree_map(add_one, list1) |
| 69 | tuple2 = mlx.utils.tree_map(add_one, tuple1) |
| 70 | vector2 = mlx.utils.tree_map(add_one, vector1) |
| 71 | params2 = mlx.utils.tree_map(add_one, params1) |
| 72 | dict2 = mlx.utils.tree_map(add_one, dict1) |
| 73 | |
| 74 | self.assertTrue(isinstance(list2, list)) |
| 75 | self.assertTrue(mx.array_equal(list2[0], mx.array([1, 2]))) |
| 76 | self.assertTrue(mx.array_equal(list2[1], mx.array(3))) |
| 77 | |
| 78 | self.assertTrue(isinstance(tuple2, tuple)) |
| 79 | self.assertTrue(mx.array_equal(tuple2[0], mx.array([1, 2]))) |
| 80 | self.assertTrue(mx.array_equal(tuple2[1], mx.array(3))) |
| 81 | |
| 82 | self.assertTrue(isinstance(vector2, Vector)) |
| 83 | self.assertTrue(mx.array_equal(vector2[0], mx.array([1, 2]))) |
| 84 | self.assertTrue(mx.array_equal(vector2[1], mx.array(3))) |
| 85 | |
| 86 | self.assertTrue(isinstance(dict2, dict)) |
| 87 | self.assertTrue(mx.array_equal(dict2["m"], mx.array([1, 2]))) |
| 88 | self.assertTrue(mx.array_equal(dict2["b"], mx.array(3))) |
| 89 | |
| 90 | self.assertTrue(isinstance(params2, Params)) |
| 91 | self.assertTrue(mx.array_equal(params2.m, mx.array([1, 2]))) |
| 92 | self.assertTrue(mx.array_equal(params2.b, mx.array(3))) |
| 93 | |
| 94 | |
| 95 | if __name__ == "__main__": |