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
11 changes: 10 additions & 1 deletion python/tvm/target/detect_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,12 +23,21 @@


def _detect_metal(dev: Device) -> Target:
f_get_target_property = get_global_func("device_api.metal.get_target_property")
return Target(
{
"kind": "metal",
"max_shared_memory_per_block": 32768,
"max_shared_memory_per_block": dev.max_shared_memory_per_block,
"max_threads_per_block": dev.max_threads_per_block,
"thread_warp_size": dev.warp_size,
"metal_language_version": f_get_target_property(dev, "metal_language_version"),
"supports_bfloat16": f_get_target_property(dev, "supports_bfloat16"),
"supports_simdgroup_permute": f_get_target_property(dev, "supports_simdgroup_permute"),
"supports_simdgroup_reduction": f_get_target_property(
dev, "supports_simdgroup_reduction"
),
"supports_simdgroup_matrix": f_get_target_property(dev, "supports_simdgroup_matrix"),
"supports_metal4": f_get_target_property(dev, "supports_metal4"),
}
)

Expand Down
38 changes: 38 additions & 0 deletions src/runtime/metal/metal_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -305,6 +305,35 @@ class MetalRawStream final : public Stream {
};


/*!
* \brief Device facts that decide which "metal" target attributes a device supports.
*
* Queried once per device when the workspace initializes. Python target
* detection reads them through "device_api.metal.get_target_property" to
* populate the matching attributes of the "metal" target kind, the same way
* the Vulkan device API reports its VulkanDeviceProperties.
*/
struct MetalDeviceProperties {
/*!
* \brief Highest Metal Shading Language version the device can compile,
* encoded as major * 10 + minor (for example 31 for MSL 3.1).
*/
int metal_language_version;
bool supports_bfloat16;
bool supports_simdgroup_permute;
bool supports_simdgroup_reduction;
bool supports_simdgroup_matrix;
bool supports_metal4;
};

/*!
* \brief Map a metal_language_version attribute value to the MTLLanguageVersion enum.
* \param version The version encoded as major * 10 + minor.
* \return The matching enum value.
* \note Throws when this build's SDK does not know the requested version.
*/
MTLLanguageVersion MetalLanguageVersionFromNumber(int version);

/*!
* \brief Process global Metal workspace.
*/
Expand All @@ -314,6 +343,8 @@ class MetalWorkspace final : public DeviceAPI {
std::vector<id<MTLDevice>> devices;
// Warp size constant
std::vector<int> warp_size;
// Target-relevant properties of each device, parallel to `devices`.
std::vector<MetalDeviceProperties> device_properties;
MetalWorkspace();
// Destructor
~MetalWorkspace();
Expand All @@ -327,6 +358,13 @@ class MetalWorkspace final : public DeviceAPI {
// override device API
void SetDevice(Device dev) final;
void GetAttr(Device dev, DeviceAttrKind kind, ffi::Any* rv) final;
/*!
* \brief Report one "metal" target attribute supported by the device.
* \param dev The device to query.
* \param property The target attribute name, e.g. "supports_bfloat16".
* \param rv The value; left unset when the property is unknown.
*/
void GetTargetProperty(Device dev, const std::string& property, ffi::Any* rv);
void* AllocDataSpace(Device dev, size_t nbytes, size_t alignment, DLDataType type_hint) final;
void FreeDataSpace(Device dev, void* ptr) final;
TVMStreamHandle CreateStream(Device dev) final;
Expand Down
133 changes: 131 additions & 2 deletions src/runtime/metal/metal_device_api.mm
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,100 @@
return inst;
}

namespace {

// Highest Metal Shading Language version the running OS can compile, encoded
// as major * 10 + minor. The SDK guards keep older toolchains building; the
// @available checks keep the answer correct on older systems.
int MaximumMetalLanguageVersion() {
int version = 23;
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 120000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 150000)
if (@available(macOS 12.0, iOS 15.0, *)) version = 24;
#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, *)) version = 30;
#endif
#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, *)) version = 31;
#endif
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 150000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 180000)
if (@available(macOS 15.0, iOS 18.0, *)) version = 32;
#endif
return version;
}

bool SupportsFamily(id<MTLDevice> device, MTLGPUFamily family) {
if (@available(macOS 10.15, iOS 13.0, *)) return [device supportsFamily:family];
return false;
}

MetalDeviceProperties QueryDeviceProperties(id<MTLDevice> device) {
// Feature availability follows the Metal feature set tables:
// https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf
const bool apple6 = SupportsFamily(device, MTLGPUFamilyApple6);
const bool apple7 = SupportsFamily(device, MTLGPUFamilyApple7);
const bool mac2 = SupportsFamily(device, MTLGPUFamilyMac2);
bool metal4 = false;
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 260000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 260000)
if (@available(macOS 26.0, iOS 26.0, *)) {
metal4 = [device supportsFamily:MTLGPUFamilyMetal4];
}
#endif
// MSL 4.0 is only offered to devices in the Metal 4 family; every other
// device stops at the highest 3.x version the OS provides.
const int language_version = metal4 ? 40 : MaximumMetalLanguageVersion();
MetalDeviceProperties properties;
properties.metal_language_version = language_version;
properties.supports_bfloat16 = language_version >= 31 && (apple6 || mac2);
properties.supports_simdgroup_permute = apple6 || mac2;
properties.supports_simdgroup_reduction = apple7 || mac2;
properties.supports_simdgroup_matrix = apple7;
properties.supports_metal4 = metal4;
return properties;
}

} // namespace

MTLLanguageVersion MetalLanguageVersionFromNumber(int version) {
switch (version) {
case 23:
return MTLLanguageVersion2_3;
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 120000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 150000)
case 24:
return MTLLanguageVersion2_4;
#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)
case 30:
return MTLLanguageVersion3_0;
#endif
#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)
case 31:
return MTLLanguageVersion3_1;
#endif
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 150000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 180000)
case 32:
return MTLLanguageVersion3_2;
#endif
#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 260000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 260000)
case 40:
return MTLLanguageVersion4_0;
#endif
default:
TVM_FFI_THROW(RuntimeError) << "Metal language version " << version
<< " is not available in this build";
}
}

void MetalWorkspace::GetAttr(Device dev, DeviceAttrKind kind, ffi::Any* rv) {
AUTORELEASEPOOL {
size_t index = static_cast<size_t>(dev.device_id);
Expand All @@ -66,8 +160,10 @@
#endif
break;
}
case kMaxSharedMemoryPerBlock:
return;
case kMaxSharedMemoryPerBlock: {
*rv = static_cast<int64_t>([devices[dev.device_id] maxThreadgroupMemoryLength]);
break;
}
case kComputeVersion:
return;
case kDeviceName:
Expand Down Expand Up @@ -102,6 +198,31 @@
};
}

void MetalWorkspace::GetTargetProperty(Device dev, const std::string& property, ffi::Any* rv) {
size_t index = static_cast<size_t>(dev.device_id);
TVM_FFI_ICHECK_LT(index, device_properties.size()) << "Invalid device id " << index;
const MetalDeviceProperties& properties = device_properties[index];

if (property == "metal_language_version") {
*rv = int64_t(properties.metal_language_version);
}
if (property == "supports_bfloat16") {
*rv = properties.supports_bfloat16;
}
if (property == "supports_simdgroup_permute") {
*rv = properties.supports_simdgroup_permute;
}
if (property == "supports_simdgroup_reduction") {
*rv = properties.supports_simdgroup_reduction;
}
if (property == "supports_simdgroup_matrix") {
*rv = properties.supports_simdgroup_matrix;
}
if (property == "supports_metal4") {
*rv = properties.supports_metal4;
}
}

static const char* kDummyKernel = R"A0B0(
using namespace metal;
// Simple copy kernel
Expand Down Expand Up @@ -163,13 +284,15 @@ int GetWarpSize(id<MTLDevice> dev) {
// on iPhone
id<MTLDevice> d = MTLCreateSystemDefaultDevice();
devices.push_back(d);
device_properties.push_back(QueryDeviceProperties(d));
#else
NSArray<id<MTLDevice> >* devs = MTLCopyAllDevices();
for (size_t i = 0; i < devs.count; ++i) {
id<MTLDevice> d = [devs objectAtIndex:i];
devices.push_back(d);
DLOG(INFO) << "Intializing Metal device " << i << ", name=" << [d.name UTF8String];
warp_size.push_back(GetWarpSize(d));
device_properties.push_back(QueryDeviceProperties(d));
}
#endif
this->ReinitializeDefaultStreams();
Expand Down Expand Up @@ -400,6 +523,12 @@ int GetWarpSize(id<MTLDevice> dev) {
DeviceAPI* ptr = MetalWorkspace::Global();
*rv = static_cast<void*>(ptr);
})
.def("device_api.metal.get_target_property",
[](Device dev, const std::string& property) {
ffi::Any rv;
MetalWorkspace::Global()->GetTargetProperty(dev, property, &rv);
return rv;
})
.def("metal.ResetGlobalState",
[]() { MetalWorkspace::Global()->ReinitializeDefaultStreams(); })
.def("metal.GetProfileCounters",
Expand Down
58 changes: 27 additions & 31 deletions src/runtime/metal/metal_module.mm
Original file line number Diff line number Diff line change
Expand Up @@ -45,26 +45,12 @@
#include "metal_common.h"
#include "tvm/runtime/device_api.h"

#if (defined(__MAC_OS_X_VERSION_MAX_ALLOWED) && __MAC_OS_X_VERSION_MAX_ALLOWED >= 260000) || \
(defined(__IPHONE_OS_VERSION_MAX_ALLOWED) && __IPHONE_OS_VERSION_MAX_ALLOWED >= 260000)
#define TVM_METAL_HAS_MSL_4_0 1
#endif

namespace tvm {
namespace runtime {

/*! \brief Maximum number of GPU supported in MetalModule. */
static constexpr const int kMetalMaxNumDevice = 32;

static bool MetalDeviceSupportsMetal4(id<MTLDevice> device) {
#if defined(TVM_METAL_HAS_MSL_4_0)
if (@available(macOS 26.0, iOS 26.0, *)) {
return [device supportsFamily:MTLGPUFamilyMetal4];
}
#endif
return false;
}

// Module to support thread-safe multi-GPU execution.
// The runtime will contain a per-device module table
// The modules will be lazily loaded
Expand All @@ -74,14 +60,17 @@ static bool MetalDeviceSupportsMetal4(id<MTLDevice> device) {
// src/target/metal/metal_fallback_module.h. The per-kernel `smap`
// payload is Map<String, Bytes> regardless of whether the format is
// text MSL ("metal") or compiled metallib ("metallib") — text vs binary
// distinction lives in `fmt`.
// distinction lives in `fmt`. `metal_language_version` is the MSL version
// codegen assumed (major * 10 + minor); fmt="metal" sources are compiled
// with exactly that version on every device.
MetalModuleNode(ffi::Map<ffi::String, ffi::Bytes> smap, ffi::String fmt,
ffi::Map<ffi::String, FunctionInfo> fmap,
ffi::Map<ffi::String, ffi::String> source)
ffi::Map<ffi::String, ffi::String> source, int metal_language_version)
: smap_(std::move(smap)),
fmt_(std::move(fmt)),
fmap_(std::move(fmap)),
source_(std::move(source)) {}
source_(std::move(source)),
metal_language_version_(metal_language_version) {}

const char* kind() const final { return "metal"; }

Expand All @@ -93,14 +82,16 @@ int GetPropertyMask() const final {
ffi::Optional<ffi::Function> GetFunction(const ffi::String& name) final;

ffi::Bytes SaveToBytes() const final {
// 3 fields [fmt][fmap][smap]. Source map is in-memory inspection only
// and is NEVER serialized — matches the cross-backend rule.
// 4 fields [fmt][metal_language_version][fmap][smap]. Source map is
// in-memory inspection only and is NEVER serialized — matches the
// cross-backend rule.
// MetalFallbackModuleNode::SaveToBytes (in
// src/target/metal/metal_fallback_module.cc) MUST mirror this format
// byte-for-byte; see one-way comment there.
std::string result;
support::BytesOutStream stream(&result);
stream.Write(fmt_);
stream.Write(metal_language_version_);
stream.Write(fmap_);
stream.Write(smap_);
return ffi::Bytes(std::move(result));
Expand Down Expand Up @@ -138,14 +129,13 @@ int GetPropertyMask() const final {
const ffi::Bytes& source = (*kernel).second;

if (fmt_ == "metal") {
const int device_language_version =
w->device_properties[device_id].metal_language_version;
TVM_FFI_ICHECK_LE(metal_language_version_, device_language_version)
<< "Metal module was generated for MSL " << metal_language_version_ << " but device "
<< device_id << " compiles at most MSL " << device_language_version;
MTLCompileOptions* opts = [[MTLCompileOptions alloc] init];
MTLLanguageVersion language_version = MTLLanguageVersion2_3;
#if defined(TVM_METAL_HAS_MSL_4_0)
if (MetalDeviceSupportsMetal4(w->devices[device_id])) {
language_version = MTLLanguageVersion4_0;
}
#endif
opts.languageVersion = language_version;
opts.languageVersion = metal::MetalLanguageVersionFromNumber(metal_language_version_);
opts.fastMathEnabled = YES;
// Per-kernel payload is bytes; treat as UTF-8 MSL source.
std::string source_str(source.data(), source.size());
Expand Down Expand Up @@ -211,6 +201,8 @@ int GetPropertyMask() const final {
ffi::Map<ffi::String, FunctionInfo> fmap_;
// In-memory source map for InspectSource — never serialized.
ffi::Map<ffi::String, ffi::String> source_;
// MSL version the kernels were generated for (major * 10 + minor).
int metal_language_version_;
// function information.
std::vector<DeviceEntry> finfo_;
// internal mutex when updating the module
Expand Down Expand Up @@ -330,26 +322,29 @@ void operator()(ffi::PackedArgs args, ffi::Any* rv, const ArgUnion64* pack_args)

static ffi::Module MetalModuleCreateImpl(ffi::Map<ffi::String, ffi::Bytes> smap, ffi::String fmt,
ffi::Map<ffi::String, FunctionInfo> fmap,
ffi::Map<ffi::String, ffi::String> source) {
ffi::Map<ffi::String, ffi::String> source,
int metal_language_version) {
ffi::ObjectPtr<MetalModuleNode> n;
AUTORELEASEPOOL {
n = ffi::make_object<MetalModuleNode>(std::move(smap), std::move(fmt), std::move(fmap),
std::move(source));
std::move(source), metal_language_version);
};
return ffi::Module(n);
}

static ffi::Module MetalModuleLoadFromBytes(const ffi::Bytes& bytes) {
support::BytesInStream stream(bytes);
ffi::String fmt;
int metal_language_version;
ffi::Map<ffi::String, FunctionInfo> fmap;
ffi::Map<ffi::String, ffi::Bytes> smap;
stream.Read(&fmt);
TVM_FFI_ICHECK(stream.Read(&metal_language_version));
TVM_FFI_ICHECK(stream.Read(&fmap));
stream.Read(&smap);
// Source map is not serialized — reconstructed empty on load.
return MetalModuleCreateImpl(std::move(smap), std::move(fmt), std::move(fmap),
ffi::Map<ffi::String, ffi::String>());
ffi::Map<ffi::String, ffi::String>(), metal_language_version);
}

void SetMetalStream(TVMStreamHandle stream) {
Expand All @@ -371,9 +366,10 @@ void SetMetalStream(TVMStreamHandle stream) {
.def("ffi.Module.load_from_bytes.metal", MetalModuleLoadFromBytes)
.def("ffi.Module.create.metal",
[](ffi::Map<ffi::String, ffi::Bytes> smap, ffi::String fmt,
ffi::Map<ffi::String, FunctionInfo> fmap, ffi::Map<ffi::String, ffi::String> source) {
ffi::Map<ffi::String, FunctionInfo> fmap, ffi::Map<ffi::String, ffi::String> source,
int metal_language_version) {
return MetalModuleCreateImpl(std::move(smap), std::move(fmt), std::move(fmap),
std::move(source));
std::move(source), metal_language_version);
})
.def("metal.SetStream", SetMetalStream);
}
Expand Down
Loading