Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion src/backend/metal/runtime/metal_module.mm
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,25 @@ static bool MetalDeviceSupportsMetal4(id<MTLDevice> device) {
return false;
}

// Highest MSL version below 4.0 that the running OS compiles. MSL 3.0
// (macOS 13 / iOS 16) adds device atomic<float>; 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
Expand Down Expand Up @@ -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;
Expand Down
46 changes: 46 additions & 0 deletions tests/python/codegen/test_target_codegen_metal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading