Skip to content

Commit 8db6939

Browse files
committed
Fix cross-platform CI dispatch checks
1 parent 229c149 commit 8db6939

3 files changed

Lines changed: 46 additions & 5 deletions

File tree

evals/mlx_cases.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -332,6 +332,10 @@ def _model_dispatch(*, supports_kernels: bool) -> str:
332332
return "general-mlx"
333333
if not supports_kernels:
334334
return "general-mlx (model has no kernel toggle)"
335+
from e3nn_mlx.compat import mlx_metal_available
336+
337+
if not mlx_metal_available():
338+
return "general-mlx (kernel fallback)"
335339
return "mixed-model-kernels"
336340

337341

tests/test_evals.py

Lines changed: 34 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -163,6 +163,7 @@ def test_backend_builders_match_the_documented_case_set() -> None:
163163
@pytest.mark.mlx
164164
def test_model_benchmarks_report_kernel_scope() -> None:
165165
from evals import mlx_cases
166+
from e3nn_mlx.compat import mlx_metal_available
166167

167168
config = {
168169
"nodes": 8,
@@ -178,13 +179,19 @@ def test_model_benchmarks_report_kernel_scope() -> None:
178179
mlx_cases.configure(use_custom_kernels=False)
179180

180181
assert gate.dispatch == "general-mlx (model has no kernel toggle)"
181-
assert simple.dispatch == "mixed-model-kernels"
182-
assert attributed.dispatch == "mixed-model-kernels"
182+
expected = (
183+
"mixed-model-kernels"
184+
if mlx_metal_available()
185+
else "general-mlx (kernel fallback)"
186+
)
187+
assert simple.dispatch == expected
188+
assert attributed.dispatch == expected
183189

184190

185191
@pytest.mark.mlx
186192
def test_mlx_benchmark_reports_actual_tensor_product_dispatch() -> None:
187193
from evals import mlx_cases
194+
from e3nn_mlx.compat import mlx_metal_available
188195

189196
mlx_cases.configure(use_custom_kernels=True)
190197
small = mlx_cases.build_fully_connected_tensor_product(
@@ -198,6 +205,30 @@ def test_mlx_benchmark_reports_actual_tensor_product_dispatch() -> None:
198205
{"items": 16, "mul": 8, "lmax": 2}
199206
)
200207

201-
assert small.dispatch == "metal-scalar-paths"
208+
expected = (
209+
"metal-scalar-paths"
210+
if mlx_metal_available()
211+
else "general-mlx (kernel fallback)"
212+
)
213+
assert small.dispatch == expected
202214
assert dense.dispatch == "general-mlx (kernel fallback)"
203215
assert general.dispatch == "general-mlx"
216+
217+
218+
@pytest.mark.mlx
219+
def test_mlx_benchmark_reports_non_metal_kernel_fallback(monkeypatch) -> None:
220+
import e3nn_mlx.compat as compat
221+
import e3nn_mlx.ops_tp as tp_module
222+
from evals import mlx_cases
223+
224+
monkeypatch.setattr(compat, "mlx_metal_available", lambda: False)
225+
monkeypatch.setattr(tp_module, "mlx_metal_available", lambda: False)
226+
mlx_cases.configure(use_custom_kernels=True)
227+
task = mlx_cases.build_fully_connected_tensor_product(
228+
{"items": 16, "mul": 8, "lmax": 2}
229+
)
230+
model_dispatch = mlx_cases._model_dispatch(supports_kernels=True)
231+
mlx_cases.configure(use_custom_kernels=False)
232+
233+
assert task.dispatch == "general-mlx (kernel fallback)"
234+
assert model_dispatch == "general-mlx (kernel fallback)"

tests/test_v2106_convolution.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -154,11 +154,17 @@ def test_v2106_convolution_compilation_and_cached_reuse() -> None:
154154
expected = module.forward_arrays(*values)
155155
mx.eval(expected)
156156
compiled = mx.compile(module.forward_arrays)
157-
assert _max_abs(compiled(*values) - expected) < 5e-5
157+
actual = compiled(*values)
158+
mx.eval(actual)
159+
assert bool(mx.allclose(actual, expected, atol=5e-5, rtol=2e-6))
158160
changed = (values[0] * 0.9, values[1], *values[2:])
159161
changed_expected = module.forward_arrays(*changed)
160162
mx.eval(changed_expected)
161-
assert _max_abs(compiled(*changed) - changed_expected) < 5e-5
163+
changed_actual = compiled(*changed)
164+
mx.eval(changed_actual)
165+
assert bool(
166+
mx.allclose(changed_actual, changed_expected, atol=5e-5, rtol=2e-6)
167+
)
162168

163169

164170
@pytest.mark.mlx

0 commit comments

Comments
 (0)