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
2 changes: 2 additions & 0 deletions ggml/src/ggml-metal/ggml-metal-device.h
Original file line number Diff line number Diff line change
Expand Up @@ -246,6 +246,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];
Expand Down
68 changes: 40 additions & 28 deletions ggml/src/ggml-metal/ggml-metal-device.m
Original file line number Diff line number Diff line change
Expand Up @@ -927,6 +927,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;
Expand All @@ -938,39 +966,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));

Expand Down
2 changes: 1 addition & 1 deletion ggml/src/ggml-metal/ggml-metal-tuning.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ fa_vec_cfg_t fa_vec_baseline_cfg(int dk, int dv) {
return { 1, (int8_t) fa_vec_baseline_ne(dk, dv) };
}

// Generated by `test-backend-ops tune --tune-perf`; do not hand-edit.
// Generated by `ggml-metal-tuning fa-vec`; do not hand-edit.
// One row per kept bucket, plus per-(dtype,dk,dv) ne11-collapsed domain defaults
// (ne11_b = FA_VEC_NE11_DEFAULT, ne01_b = domain). To retune or add a device, re-run the
// sweep and paste its block. See ggml-metal-tuning.h for the row/lookup semantics.
Expand Down
18 changes: 16 additions & 2 deletions ggml/src/ggml-metal/ggml-metal-tuning.h
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
#include "ggml.h"

#include <cstdint>
#include <vector>

namespace ggml_metal_tuning {

Expand All @@ -18,8 +19,8 @@ int fa_vec_ne11_bucket(int64_t ne11);
int fa_vec_ne01_bucket(int64_t ne01);

// NE baked into each (dk,dv) baseline instantiation in kernels/fa.metal.
// Hand-maintained mirror; keep in sync with those instantiations (run_fa_vec_tune_check
// exercises every (Q,NE), so a missing instantiation surfaces there).
// Hand-maintained mirror; keep in sync with those instantiations.
// The Metal test slice covers every legal config for dk=128 and dk=576.
int fa_vec_baseline_ne(int dk, int dv);

// Tuned table has two row kinds. Exact rows key a (ne11_b, ne01_b) bucket. Default rows
Expand Down Expand Up @@ -51,6 +52,19 @@ struct fa_vec_entry_t {
fa_vec_cfg_t cfg;
};

// legal NE values for a (dk,dv): NL = 32/NE, require (dk/4)%NL==0 && (dv/4)%NL==0.
// single source shared by the offline tuner and test-backend-ops.
inline std::vector<int> fa_vec_legal_ne(int dk, int dv) {
std::vector<int> r;
for (int ne : { 1, 2, 4 }) {
const int nl = 32 / ne;
if ((dk / 4) % nl == 0 && (dv / 4) % nl == 0) {
r.push_back(ne);
}
}
return r;
}

// test/tune-only override; when set, fa_vec_pick returns it directly.
void fa_vec_set_override(fa_vec_cfg_t cfg);
void fa_vec_clear_override();
Expand Down
9 changes: 9 additions & 0 deletions ggml/src/ggml-metal/ggml-metal.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;

Expand Down
Loading
Loading