From 178477c0a6e96afcf4a062ccfd7b429caf5ff1ba Mon Sep 17 00:00:00 2001 From: baigengyuan Date: Sun, 4 Oct 2026 00:18:18 -0400 Subject: [PATCH] [Metal] Compile MSL source with the highest language version the OS supports The Metal runtime compiles every textual MSL module with MTLLanguageVersion2_3 unless the device supports Metal 4. MSL 2.3 has no bfloat type (added in MSL 3.1) and no device atomic (added in MSL 3.0), so any kernel that uses them fails to compile at load time on macOS 13-15, even though the OS supports a newer language version. The Metal codegen prints bfloat16 as `bfloat` (including `simdgroup_bfloat8x8` fragments), and MSL supplied through tvm_callback_metal_compile or by downstream projects hits the same limit. Select the highest MSL version below 4.0 that the running OS supports: MSL 3.1 on macOS 14 / iOS 17 and later, MSL 3.0 on macOS 13 / iOS 16, and MSL 2.3 otherwise. Each newer version is guarded by both the SDK version macros and @available, so builds against older SDKs and older deployment targets keep the current behavior. The MSL 4.0 selection for Metal 4 devices is unchanged. Add a test that compiles a kernel through tvm_callback_metal_compile using `bfloat` and __METAL_VERSION__, and checks that the runtime compiled it as MSL 3.1 or newer on macOS 14+. --- src/backend/metal/runtime/metal_module.mm | 21 ++++++++- .../codegen/test_target_codegen_metal.py | 46 +++++++++++++++++++ 2 files changed, 66 insertions(+), 1 deletion(-) diff --git a/src/backend/metal/runtime/metal_module.mm b/src/backend/metal/runtime/metal_module.mm index a283c27fd8c2..d062d6aacd4c 100644 --- a/src/backend/metal/runtime/metal_module.mm +++ b/src/backend/metal/runtime/metal_module.mm @@ -64,6 +64,25 @@ static bool MetalDeviceSupportsMetal4(id device) { return false; } +// Highest MSL version below 4.0 that the running OS compiles. MSL 3.0 +// (macOS 13 / iOS 16) adds device atomic; MSL 3.1 (macOS 14 / iOS 17) +// adds bfloat, which the Metal codegen emits for bfloat16. +static MTLLanguageVersion MetalDefaultLanguageVersion() { +#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 140000) || \ + (defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 170000) + if (@available(macOS 14.0, iOS 17.0, *)) { + return MTLLanguageVersion3_1; + } +#endif +#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 130000) || \ + (defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 160000) + if (@available(macOS 13.0, iOS 16.0, *)) { + return MTLLanguageVersion3_0; + } +#endif + return MTLLanguageVersion2_3; +} + // Module to support thread-safe multi-GPU execution. // The runtime will contain a per-device module table // The modules will be lazily loaded @@ -138,7 +157,7 @@ int GetPropertyMask() const final { if (fmt_ == "metal") { MTLCompileOptions* opts = [[MTLCompileOptions alloc] init]; - MTLLanguageVersion language_version = MTLLanguageVersion2_3; + MTLLanguageVersion language_version = MetalDefaultLanguageVersion(); #if defined(TVM_METAL_HAS_MSL_4_0) if (MetalDeviceSupportsMetal4(w->devices[device_id])) { language_version = MTLLanguageVersion4_0; diff --git a/tests/python/codegen/test_target_codegen_metal.py b/tests/python/codegen/test_target_codegen_metal.py index 36c053e3a99b..687782af2cbd 100644 --- a/tests/python/codegen/test_target_codegen_metal.py +++ b/tests/python/codegen/test_target_codegen_metal.py @@ -14,6 +14,8 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +import platform + import numpy as np import pytest import tvm_ffi @@ -651,5 +653,49 @@ def run_and_check(): tvm.testing.run_with_gpu_lock(run_and_check) +def _macos_major_version(): + version = platform.mac_ver()[0] + return int(version.split(".")[0]) if version else 0 + + +@pytest.mark.gpu +@pytest.mark.skipif(not env.has_metal(), reason="need metal") +@pytest.mark.skipif(_macos_major_version() < 14, reason="MSL 3.1 requires macOS 14") +def test_metal_source_compiled_with_msl_3_1(): + """Textual MSL is compiled with MSL 3.1 or newer, so `bfloat` is available.""" + n = 32 + + @I.ir_module + class Module: + @T.prim_func + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): + T.func_attr({"tirx.noalias": True}) + for i in T.thread_binding(n, thread="threadIdx.x"): + B[i] = A[i] + T.float32(1.0) + + def msl_version_callback(src, target): + # Report the MSL version the runtime compiled with, and use `bfloat`, + # which MSL 2.3 does not define. + new_src = src.replace("1.000000e+00f", "((float)__METAL_VERSION__ + (float)bfloat(1.0f))") + assert new_src != src + return (new_src, "metal") + + tvm.register_global_func("tvm_callback_metal_compile", msl_version_callback, override=True) + try: + f = tvm.compile(Module, target="metal") + finally: + tvm_ffi.registry.remove_global_func("tvm_callback_metal_compile") + + def run_and_check(): + dev = tvm.metal(0) + a = tvm.runtime.tensor(np.zeros(n, dtype="float32"), dev) + b = tvm.runtime.empty((n,), "float32", dev) + f(a, b) + # __METAL_VERSION__ is 310 for MSL 3.1. + assert int(b.numpy()[0]) >= 311 + + tvm.testing.run_with_gpu_lock(run_and_check) + + if __name__ == "__main__": tvm.testing.main()