test: probe production split-k Metal kernel

This commit is contained in:
Jourloy 2026-07-17 09:27:07 +03:00
parent b50b5c4fc5
commit 83cf977a8f

View File

@ -31,13 +31,14 @@ def __detect_defaults():
def __flex_gemm_works_on_mps(): def __flex_gemm_works_on_mps():
"""Probe flex_gemm with tiny MPS convs covering both IMPLICIT_GEMM and """Probe dense, masked, and production split-k flex_gemm kernels on MPS.
MASKED_IMPLICIT_GEMM. If the install pre-dates the device-routing fix
(or pre-dates the real masked kernel in round 2), one of these returns If the install pre-dates the device-routing fix (or a real masked kernel),
a CPU tensor fall back to the pure-PyTorch backend rather than one of these returns a CPU tensor or raises. Fall back to pure PyTorch
crashing inside the model on the first LayerNorm. Build tensors on CPU rather than crashing inside the model on the first LayerNorm. Build tensors
and move to MPS because some PyTorch builds lack int/fp16 torch.zeros on CPU and move to MPS because some PyTorch builds lack int/fp16 creation
kernels on MPS.""" kernels on MPS.
"""
try: try:
import os import os
if os.environ.get('TRELLIS_DISABLE_METAL', '0') == '1': if os.environ.get('TRELLIS_DISABLE_METAL', '0') == '1':
@ -56,10 +57,15 @@ def __flex_gemm_works_on_mps():
weight = torch.zeros((4, 1, 1, 1, 4), dtype=torch.float16).to('mps') weight = torch.zeros((4, 1, 1, 1, 4), dtype=torch.float16).to('mps')
shape = torch.Size([1, 4, 1, 1, 1]) shape = torch.Size([1, 4, 1, 1, 1])
for algo in (Algorithm.IMPLICIT_GEMM, Algorithm.MASKED_IMPLICIT_GEMM): for algo in (
Algorithm.IMPLICIT_GEMM,
Algorithm.MASKED_IMPLICIT_GEMM,
Algorithm.MASKED_IMPLICIT_GEMM_SPLITK,
):
set_algorithm(algo) set_algorithm(algo)
out, _ = sparse_submanifold_conv3d(feats, coords, shape, weight) out, _ = sparse_submanifold_conv3d(feats, coords, shape, weight)
if out.device.type != 'mps': torch.mps.synchronize()
if out.device.type != 'mps' or not bool(torch.isfinite(out).all().item()):
return False return False
return True return True
except Exception: except Exception: