@@ -163,6 +163,7 @@ def test_backend_builders_match_the_documented_case_set() -> None:
163163@pytest .mark .mlx
164164def 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
186192def 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)"
0 commit comments