metal : per-device tuned (Q, NE) for flash-attn vec (#26570)
* metal : per-device tuned (Q, NE) for flash-attn vec (#25750) * rebase Q-generic FA vec body from 01dc93607 (#23114) * add 53 f16 (Q,NE) flash-attn vec instantiations (vec 80 -> 133) * add FA vec (Q,NE) tuning table + dispatch wiring + SMEM cap fallback * add FA vec (Q,NE) perf sweep * fill tuning result * fold family table into a per-family representative SKU * refactor tuning result format * extend FA vec tuning to quantized KV caches * sync fa vec tuner bucketing with runtime, use pointwise tuning regret * update tuned table * format and cleanup * prefix fa_vec tuning procs with ggml_backend_metal_tuning_, drop unused fa_vec_override_active * add device id -> token lookup for the offline tuning tool * add ggml-metal-tuning skeleton * add op-agnostic perf cell + median timing for the tuner * add FA-vec graph build + tensor init to the tuner * tools : add FA-vec (Q,NE) sweep, compression and table emit * cool down and re-measure the dirty window on thermal drift * test-backend-ops : replace the FA vec tune mode with a bounded (Q,NE) slice * tools : document the Metal tuner, point the table comment at it * abort on unknown KV type, single-source fa_vec_legal_ne * cleanup * honor -o in the FA vec (Q,NE) slice * retune FA-vec (Q, NE) under a pointwise no-harm gate * cont : add fa-vec tunings for M1 Pro, M2 Ultra, M5 Max --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
co-authored by
Georgi Gerganov
parent
b615f5b4bd
commit
f280b26983
@@ -11,6 +11,7 @@ ggml_add_backend_library(ggml-metal
|
||||
ggml-metal-common.cpp
|
||||
ggml-metal-context.m
|
||||
ggml-metal-ops.cpp
|
||||
ggml-metal-tuning.cpp
|
||||
)
|
||||
|
||||
target_link_libraries(ggml-metal PRIVATE
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
#include "ggml-metal-device.h"
|
||||
|
||||
#include "ggml-metal-impl.h"
|
||||
#include "ggml-metal-tuning.h"
|
||||
|
||||
#include "ggml-impl.h"
|
||||
|
||||
@@ -1544,6 +1545,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_v
|
||||
bool has_bias,
|
||||
bool has_scap,
|
||||
bool has_kvpad,
|
||||
int32_t nqpsg,
|
||||
int32_t ne,
|
||||
int32_t nsg,
|
||||
int32_t nwg,
|
||||
bool use_kv_f16,
|
||||
@@ -1559,11 +1562,17 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_v
|
||||
|
||||
const char * type = use_kv_f16 ? "f16" : ggml_type_name(op->src[1]->type);
|
||||
|
||||
snprintf(base, 256, "kernel_%s_%s_dk%d_dv%d",
|
||||
char qne_suffix[16] = {0};
|
||||
if (!(nqpsg == 1 && ne == ggml_metal_tuning::fa_vec_baseline_ne(dk, dv))) {
|
||||
snprintf(qne_suffix, sizeof(qne_suffix), "_q%d_ne%d", nqpsg, ne);
|
||||
}
|
||||
|
||||
snprintf(base, 256, "kernel_%s_%s_dk%d_dv%d%s",
|
||||
"flash_attn_ext_vec",
|
||||
type,
|
||||
dk,
|
||||
dv);
|
||||
dv,
|
||||
qne_suffix);
|
||||
|
||||
snprintf(name, 256, "%s_mask=%d_sink=%d_bias=%d_scap=%d_kvpad=%d_ns10=%d_ns20=%d_nsg=%d_nwg=%d",
|
||||
base,
|
||||
|
||||
@@ -207,6 +207,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_att
|
||||
bool has_bias,
|
||||
bool has_scap,
|
||||
bool has_kvpad,
|
||||
int32_t nqpsg,
|
||||
int32_t ne,
|
||||
int32_t nsg,
|
||||
int32_t nwg,
|
||||
bool use_kv_f16,
|
||||
@@ -257,6 +259,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;
|
||||
int device_phys;
|
||||
@@ -279,6 +283,7 @@ struct ggml_metal_device_props {
|
||||
bool supports_gpu_family_apple7;
|
||||
|
||||
enum ggml_metal_device_id device_id;
|
||||
int gpu_family;
|
||||
|
||||
int op_offload_min_batch_size;
|
||||
};
|
||||
|
||||
@@ -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, int n_devices) {
|
||||
ggml_metal_device_t dev = calloc(1, sizeof(struct ggml_metal_device));
|
||||
|
||||
@@ -1220,7 +1232,8 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) {
|
||||
{
|
||||
for (int i = MTLGPUFamilyApple1 + 20; i >= MTLGPUFamilyApple1; --i) {
|
||||
if ([dev->mtl_device supportsFamily:i]) {
|
||||
GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, i - (int) MTLGPUFamilyApple1 + 1, i);
|
||||
dev->props.gpu_family = i - (int) MTLGPUFamilyApple1 + 1;
|
||||
GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, dev->props.gpu_family, i);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
#include "ggml-metal-impl.h"
|
||||
#include "ggml-metal-common.h"
|
||||
#include "ggml-metal-device.h"
|
||||
#include "ggml-metal-tuning.h"
|
||||
|
||||
#include <cassert>
|
||||
#include <algorithm>
|
||||
@@ -3346,12 +3347,18 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
|
||||
#undef FATTN_SMEM
|
||||
} else {
|
||||
// half4x4 kernel
|
||||
const int nqptg = OP_FLASH_ATTN_EXT_VEC_NQPSG; // queries per threadgroup
|
||||
auto cfg = ggml_metal_tuning::fa_vec_pick(
|
||||
props_dev->device_id,
|
||||
props_dev->gpu_family,
|
||||
(int) op->src[1]->type,
|
||||
(int) ne00, (int) ne20, // dk, dv (ne00 == dk for FA)
|
||||
ne11, ne01);
|
||||
int nqptg = cfg.Q; // queries per threadgroup
|
||||
const int ncpsg = OP_FLASH_ATTN_EXT_VEC_NCPSG; // cache values per simdgroup !! sync with kernel template arguments !!
|
||||
const int nhptg = 1; // heads per threadgroup
|
||||
|
||||
GGML_ASSERT(nqptg <= 32);
|
||||
GGML_ASSERT(nqptg % 1 == 0);
|
||||
GGML_ASSERT(nqptg == 1 || nqptg == 2 || nqptg == 4); // only instantiated Q values
|
||||
GGML_ASSERT(ncpsg % 32 == 0);
|
||||
|
||||
bool need_sync = false;
|
||||
@@ -3410,7 +3417,7 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
|
||||
// ne20*(nsg)
|
||||
// each simdgroup has a full f32 head vector in shared mem to accumulate results
|
||||
//
|
||||
#define FATTN_SMEM(nsg) (GGML_PAD(((GGML_PAD(ne00, 128) + 4*ncpsg + 2*GGML_PAD(ne20, 128))*(nsg))*(sizeof(float)/2), 16))
|
||||
#define FATTN_SMEM(nsg) (GGML_PAD(((GGML_PAD(ne00, 128) + 4*ncpsg + 2*GGML_PAD(ne20, 128))*(nsg)*nqptg)*(sizeof(float)/2), 16))
|
||||
|
||||
int64_t nsg = 1;
|
||||
|
||||
@@ -3430,6 +3437,12 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
|
||||
}
|
||||
}
|
||||
|
||||
// fall back to baseline (Q=1) if the tuned config exceeds threadgroup memory
|
||||
if ((size_t) FATTN_SMEM(nsg) > props_dev->max_theadgroup_memory_size) {
|
||||
cfg = ggml_metal_tuning::fa_vec_baseline_cfg((int) ne00, (int) ne20);
|
||||
nqptg = cfg.Q; // = 1
|
||||
}
|
||||
|
||||
const int32_t ns10 = nb11_attn/nb10_attn;
|
||||
const int32_t ns20 = nb21_attn/nb20_attn;
|
||||
|
||||
@@ -3468,7 +3481,7 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
|
||||
/*.logit_softcap =*/ logit_softcap,
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_flash_attn_ext_vec(lib, op, has_mask, has_sinks, has_bias, has_scap, has_kvpad, nsg, nwg, use_kv_f16, ns10, ns20);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_flash_attn_ext_vec(lib, op, has_mask, has_sinks, has_bias, has_scap, has_kvpad, nqptg, cfg.NE, nsg, nwg, use_kv_f16, ns10, ns20);
|
||||
|
||||
GGML_ASSERT(nsg*32 <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,77 @@
|
||||
#pragma once
|
||||
|
||||
#include "ggml-metal-device.h" // enum ggml_metal_device_id
|
||||
#include "ggml.h"
|
||||
|
||||
#include <cstdint>
|
||||
#include <vector>
|
||||
|
||||
namespace ggml_metal_tuning {
|
||||
|
||||
// FA vec selection buckets. ne01 (query rows) splits decode (==1) from batch (>=2), the
|
||||
// batch side refined into {2,3,4,5}: Q>1 reuses one K/V load across rows, so it only pays
|
||||
// off once ne01 aligns with Q. ne11 (KV length) is bucketed too, as the Q>1 crossover is
|
||||
// head-size dependent (small dk crosses late, large dk wins even at short KV).
|
||||
constexpr int FA_VEC_NE11_BUCKETS[] = { 1024, 4096, 16384 };
|
||||
constexpr int FA_VEC_NE01_BUCKETS[] = { 2, 3, 4, 5 };
|
||||
|
||||
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.
|
||||
// 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
|
||||
// collapse ne11 over one ne01 domain: ne11_b == FA_VEC_NE11_DEFAULT and ne01_b holds the
|
||||
// domain. fa_vec_pick tries exact bucket -> domain default -> baseline; short KV
|
||||
// (ne11 < FA_VEC_NE11_BUCKETS[0]) always uses baseline.
|
||||
constexpr int8_t FA_VEC_NE11_DEFAULT = -1;
|
||||
constexpr int8_t FA_VEC_DOMAIN_DECODE = 0; // ne01 == 1
|
||||
constexpr int8_t FA_VEC_DOMAIN_BATCH = 1; // ne01 >= 2
|
||||
|
||||
struct fa_vec_key_t {
|
||||
int8_t device_id;
|
||||
int8_t dtype;
|
||||
int16_t dk;
|
||||
int16_t dv;
|
||||
int8_t ne11_b;
|
||||
int8_t ne01_b;
|
||||
};
|
||||
|
||||
static_assert(sizeof(fa_vec_key_t) == 8, "fa_vec_key_t must be tightly packed for memcmp");
|
||||
|
||||
struct fa_vec_cfg_t {
|
||||
int8_t Q;
|
||||
int8_t NE;
|
||||
};
|
||||
|
||||
struct fa_vec_entry_t {
|
||||
fa_vec_key_t key;
|
||||
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();
|
||||
fa_vec_cfg_t fa_vec_baseline_cfg(int dk, int dv);
|
||||
|
||||
// device_id selects a per-SKU row; on a miss, gpu_family (0 if unknown) maps to a representative
|
||||
// SKU and the table is retried. No match -> baseline.
|
||||
fa_vec_cfg_t fa_vec_pick(enum ggml_metal_device_id device_id, int gpu_family, int dtype, int dk, int dv, int64_t ne11, int64_t ne01);
|
||||
|
||||
} // namespace ggml_metal_tuning
|
||||
@@ -6,6 +6,7 @@
|
||||
#include "ggml-metal-device.h"
|
||||
#include "ggml-metal-context.h"
|
||||
#include "ggml-metal-ops.h"
|
||||
#include "ggml-metal-tuning.h"
|
||||
|
||||
#include <mutex>
|
||||
#include <string>
|
||||
@@ -870,10 +871,55 @@ static ggml_backend_feature * ggml_backend_metal_get_features(ggml_backend_reg_t
|
||||
GGML_UNUSED(reg);
|
||||
}
|
||||
|
||||
// test/tune-only override for the FA vec (Q, NE) selection, reached via proc_address.
|
||||
static void ggml_backend_metal_tuning_set_fa_vec_override(int Q, int NE) {
|
||||
ggml_metal_tuning::fa_vec_set_override({ (int8_t) Q, (int8_t) NE });
|
||||
}
|
||||
|
||||
static void ggml_backend_metal_tuning_clear_fa_vec_override(void) {
|
||||
ggml_metal_tuning::fa_vec_clear_override();
|
||||
}
|
||||
|
||||
static int ggml_backend_metal_tuning_fa_vec_ne11_bucket(int64_t ne11) {
|
||||
return ggml_metal_tuning::fa_vec_ne11_bucket(ne11);
|
||||
}
|
||||
|
||||
static int ggml_backend_metal_tuning_fa_vec_ne01_bucket(int64_t ne01) {
|
||||
return ggml_metal_tuning::fa_vec_ne01_bucket(ne01);
|
||||
}
|
||||
|
||||
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;
|
||||
}
|
||||
if (strcmp(name, "ggml_backend_metal_tuning_set_fa_vec_override") == 0) {
|
||||
return (void *)ggml_backend_metal_tuning_set_fa_vec_override;
|
||||
}
|
||||
if (strcmp(name, "ggml_backend_metal_tuning_clear_fa_vec_override") == 0) {
|
||||
return (void *)ggml_backend_metal_tuning_clear_fa_vec_override;
|
||||
}
|
||||
if (strcmp(name, "ggml_backend_metal_tuning_fa_vec_ne11_bucket") == 0) {
|
||||
return (void *)ggml_backend_metal_tuning_fa_vec_ne11_bucket;
|
||||
}
|
||||
if (strcmp(name, "ggml_backend_metal_tuning_fa_vec_ne01_bucket") == 0) {
|
||||
return (void *)ggml_backend_metal_tuning_fa_vec_ne01_bucket;
|
||||
}
|
||||
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;
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user