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

Method test_supported_trees

python/tests/test_tree.py:49–92  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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
95if __name__ == "__main__":

Callers

nothing calls this directly

Calls 3

ParamsClass · 0.85
VectorClass · 0.70
arrayMethod · 0.60

Tested by

no test coverage detected