diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index b7c82af74a..afde8e8bb0 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -249,6 +249,8 @@ enum ggml_metal_device_id { GGML_METAL_DEVICE_M5_ULTRA, }; +const char * ggml_metal_device_id_token(enum ggml_metal_device_id id); + struct ggml_metal_device_props { int device; char name[128]; diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 8e580a8eff..a50ccc3c53 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -962,6 +962,34 @@ void ggml_metal_rsets_free(ggml_metal_rsets_t rsets) { free(rsets); } +static const struct { + const char * name; + const char * token; + enum ggml_metal_device_id id; +} k_metal_devices[] = { +#define DEV(name, id) { name, #id, id } + DEV("M1", GGML_METAL_DEVICE_M1), + DEV("M1 Pro", GGML_METAL_DEVICE_M1_PRO), + DEV("M1 Max", GGML_METAL_DEVICE_M1_MAX), + DEV("M1 Ultra", GGML_METAL_DEVICE_M1_ULTRA), + DEV("M2", GGML_METAL_DEVICE_M2), + DEV("M2 Pro", GGML_METAL_DEVICE_M2_PRO), + DEV("M2 Max", GGML_METAL_DEVICE_M2_MAX), + DEV("M2 Ultra", GGML_METAL_DEVICE_M2_ULTRA), + DEV("M3", GGML_METAL_DEVICE_M3), + DEV("M3 Pro", GGML_METAL_DEVICE_M3_PRO), + DEV("M3 Max", GGML_METAL_DEVICE_M3_MAX), + DEV("M3 Ultra", GGML_METAL_DEVICE_M3_ULTRA), + DEV("M4", GGML_METAL_DEVICE_M4), + DEV("M4 Pro", GGML_METAL_DEVICE_M4_PRO), + DEV("M4 Max", GGML_METAL_DEVICE_M4_MAX), + DEV("M5", GGML_METAL_DEVICE_M5), + DEV("M5 Pro", GGML_METAL_DEVICE_M5_PRO), + DEV("M5 Max", GGML_METAL_DEVICE_M5_MAX), + DEV("M5 Ultra", GGML_METAL_DEVICE_M5_ULTRA), +#undef DEV +}; + static enum ggml_metal_device_id ggml_metal_device_id_parse(const char * name) { if (!name) { return GGML_METAL_DEVICE_GENERIC; @@ -973,39 +1001,23 @@ static enum ggml_metal_device_id ggml_metal_device_id_parse(const char * name) { } const char * suffix = name + sizeof(prefix) - 1; - static const struct { - const char * name; - enum ggml_metal_device_id id; - } table[] = { - {"M1", GGML_METAL_DEVICE_M1}, - {"M1 Pro", GGML_METAL_DEVICE_M1_PRO}, - {"M1 Max", GGML_METAL_DEVICE_M1_MAX}, - {"M1 Ultra", GGML_METAL_DEVICE_M1_ULTRA}, - {"M2", GGML_METAL_DEVICE_M2}, - {"M2 Pro", GGML_METAL_DEVICE_M2_PRO}, - {"M2 Max", GGML_METAL_DEVICE_M2_MAX}, - {"M2 Ultra", GGML_METAL_DEVICE_M2_ULTRA}, - {"M3", GGML_METAL_DEVICE_M3}, - {"M3 Pro", GGML_METAL_DEVICE_M3_PRO}, - {"M3 Max", GGML_METAL_DEVICE_M3_MAX}, - {"M3 Ultra", GGML_METAL_DEVICE_M3_ULTRA}, - {"M4", GGML_METAL_DEVICE_M4}, - {"M4 Pro", GGML_METAL_DEVICE_M4_PRO}, - {"M4 Max", GGML_METAL_DEVICE_M4_MAX}, - {"M5", GGML_METAL_DEVICE_M5}, - {"M5 Pro", GGML_METAL_DEVICE_M5_PRO}, - {"M5 Max", GGML_METAL_DEVICE_M5_MAX}, - {"M5 Ultra", GGML_METAL_DEVICE_M5_ULTRA}, - }; - - for (size_t i = 0; i < sizeof(table)/sizeof(table[0]); ++i) { - if (strcmp(suffix, table[i].name) == 0) { - return table[i].id; + for (size_t i = 0; i < sizeof(k_metal_devices)/sizeof(k_metal_devices[0]); ++i) { + if (strcmp(suffix, k_metal_devices[i].name) == 0) { + return k_metal_devices[i].id; } } return GGML_METAL_DEVICE_GENERIC; } +const char * ggml_metal_device_id_token(enum ggml_metal_device_id id) { + for (size_t i = 0; i < sizeof(k_metal_devices)/sizeof(k_metal_devices[0]); ++i) { + if (k_metal_devices[i].id == id) { + return k_metal_devices[i].token; + } + } + return "GGML_METAL_DEVICE_GENERIC"; +} + ggml_metal_device_t ggml_metal_device_init(int device) { ggml_metal_device_t dev = calloc(1, sizeof(struct ggml_metal_device)); diff --git a/ggml/src/ggml-metal/ggml-metal.cpp b/ggml/src/ggml-metal/ggml-metal.cpp index 32693347e6..1e40865bca 100644 --- a/ggml/src/ggml-metal/ggml-metal.cpp +++ b/ggml/src/ggml-metal/ggml-metal.cpp @@ -890,6 +890,12 @@ static int ggml_backend_metal_tuning_fa_vec_baseline_ne(int dk, int dv) { return ggml_metal_tuning::fa_vec_baseline_ne(dk, dv); } +static const char * ggml_backend_metal_tuning_device_token(ggml_backend_dev_t dev) { + ggml_metal_device_t ctx_dev = (ggml_metal_device_t)dev->context; + + return ggml_metal_device_id_token(ggml_metal_device_get_props(ctx_dev)->device_id); +} + static void * ggml_backend_metal_get_proc_address(ggml_backend_reg_t reg, const char * name) { if (strcmp(name, "ggml_backend_get_features") == 0) { return (void *)ggml_backend_metal_get_features; @@ -909,6 +915,9 @@ static void * ggml_backend_metal_get_proc_address(ggml_backend_reg_t reg, const if (strcmp(name, "ggml_backend_metal_tuning_fa_vec_baseline_ne") == 0) { return (void *)ggml_backend_metal_tuning_fa_vec_baseline_ne; } + if (strcmp(name, "ggml_backend_metal_tuning_device_token") == 0) { + return (void *)ggml_backend_metal_tuning_device_token; + } return NULL;