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()