diff --git a/src/backend/metal/runtime/metal_device_api.mm b/src/backend/metal/runtime/metal_device_api.mm index 2a44d98f109f..b8026fb84caa 100644 --- a/src/backend/metal/runtime/metal_device_api.mm +++ b/src/backend/metal/runtime/metal_device_api.mm @@ -20,6 +20,9 @@ /*! * \file metal_device_api.mm */ +#include +#include + #include #include #include @@ -184,15 +187,14 @@ int GetWarpSize(id dev) { id buf; AUTORELEASEPOOL { id dev = GetDevice(device); - // GPU memory only + // GPU memory only. Callers can override via TVM_METAL_STORAGE_MODE + // ("shared" or "managed") when Private-mode exposes a driver-level + // correctness issue on their workload — see apache/tvm#20157. MTLResourceOptions storage_mode = MTLResourceStorageModePrivate; - /* - #if TARGET_OS_IPHONE - storage_mode = MTLResourceStorageModeShared; - #else - storage_mode = MTLResourceStorageModeManaged; - #endif - */ + if (const char* env = std::getenv("TVM_METAL_STORAGE_MODE")) { + if (std::strcmp(env, "shared") == 0) storage_mode = MTLResourceStorageModeShared; + if (std::strcmp(env, "managed") == 0) storage_mode = MTLResourceStorageModeManaged; + } buf = [dev newBufferWithLength:nbytes options:storage_mode]; TVM_FFI_ICHECK(buf != nil); };