Skip to content

Commit 996826e

Browse files
committed
Stabilize compiled model comparisons
1 parent f658338 commit 996826e

4 files changed

Lines changed: 22 additions & 6 deletions

File tree

tests/test_gate_points_2102_network.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -150,13 +150,18 @@ def forward(pos, node_input, node_attr, graph_batch, edge_src, edge_dst):
150150
).array
151151

152152
expected = forward(data["pos"], data["x"], data["z"], batch, edges[0], edges[1])
153+
mx.eval(expected)
153154
compiled = mx.compile(forward)
154155
actual = compiled(data["pos"], data["x"], data["z"], batch, edges[0], edges[1])
155156
assert _max_abs(actual - expected) < 8e-5
156157
changed_x = data["x"] * 0.8
158+
changed_expected = forward(
159+
data["pos"], changed_x, data["z"], batch, edges[0], edges[1]
160+
)
161+
mx.eval(changed_expected)
157162
assert _max_abs(
158163
compiled(data["pos"], changed_x, data["z"], batch, edges[0], edges[1])
159-
- forward(data["pos"], changed_x, data["z"], batch, edges[0], edges[1])
164+
- changed_expected
160165
) < 8e-5
161166

162167

tests/test_v2106_convolution.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -105,10 +105,14 @@ def test_v2106_convolution_compilation_and_cached_reuse() -> None:
105105
module = _module()
106106
_activate_alpha(module)
107107
values = _inputs(module)
108+
expected = module.forward_arrays(*values)
109+
mx.eval(expected)
108110
compiled = mx.compile(module.forward_arrays)
109-
assert _max_abs(compiled(*values) - module.forward_arrays(*values)) < 5e-5
111+
assert _max_abs(compiled(*values) - expected) < 5e-5
110112
changed = (values[0] * 0.9, values[1], *values[2:])
111-
assert _max_abs(compiled(*changed) - module.forward_arrays(*changed)) < 5e-5
113+
changed_expected = module.forward_arrays(*changed)
114+
mx.eval(changed_expected)
115+
assert _max_abs(compiled(*changed) - changed_expected) < 5e-5
112116

113117

114118
@pytest.mark.mlx

tests/test_v2106_message_passing.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -104,10 +104,14 @@ def test_v2106_message_passing_compilation_and_cached_reuse() -> None:
104104
module = _module()
105105
_activate_alpha(module)
106106
values = _inputs(module)
107+
expected = module.forward_arrays(*values)
108+
mx.eval(expected)
107109
compiled = mx.compile(module.forward_arrays)
108-
assert _max_abs(compiled(*values) - module.forward_arrays(*values)) < 5e-4
110+
assert _max_abs(compiled(*values) - expected) < 5e-4
109111
changed = (values[0] * 1.1, values[1], *values[2:])
110-
assert _max_abs(compiled(*changed) - module.forward_arrays(*changed)) < 5e-4
112+
changed_expected = module.forward_arrays(*changed)
113+
mx.eval(changed_expected)
114+
assert _max_abs(compiled(*changed) - changed_expected) < 5e-4
111115

112116

113117
@pytest.mark.mlx

tests/test_v2106_networks.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -243,11 +243,14 @@ def forward(pos, node_input, node_attr, edge_attr, graph_batch, edge_src, edge_d
243243
edges[1],
244244
)
245245
expected = forward(*args)
246+
mx.eval(expected)
246247
compiled = mx.compile(forward)
247248
assert _max_abs(compiled(*args) - expected) < 2e-4
248249
changed = list(args)
249250
changed[1] = changed[1] * 0.9
250-
assert _max_abs(compiled(*changed) - forward(*changed)) < 2e-4
251+
changed_expected = forward(*changed)
252+
mx.eval(changed_expected)
253+
assert _max_abs(compiled(*changed) - changed_expected) < 2e-4
251254

252255

253256
@pytest.mark.mlx

0 commit comments

Comments
 (0)