Felipe97/llama-cpp-compiled
01.2k
1#include "fa-vec.h"2#include "ggml-backend.h"3#include "ggml.h"4 5#include <cstdio>6#include <cstdlib>7#include <cstring>8 9struct tuner_def {10 const char * name;11 bool (*run)(ggml_backend_t, ggml_backend_dev_t, const tuner_opts &);12};13 14static const tuner_def k_tuners[] = {15 { "fa-vec", tuner_fa_vec_run },16};17 18static void usage(const char * argv0) {19 printf("usage: %s <tuner> [options]\n", argv0);20 printf("\n");21 printf(" offline kernel tuner for the Metal backend: sweeps a kernel's config grid and\n");22 printf(" prints pasteable table rows for the machine it runs on. never a pass/fail test.\n");23 printf("\n");24 printf(" tuners:\n");25 printf(" fa-vec flash-attn vec (Q,NE) for ggml-metal-tuning.cpp\n");26 printf("\n");27 printf(" options:\n");28 printf(" -b <name> backend device (default: first Metal device)\n");29 printf(" --dtype <list> restrict KV dtypes, e.g. f16,q4_0 (default: all)\n");30 printf(" --dk <list> restrict head sizes, e.g. 128,192 (default: all)\n");31 printf(" --reps <n> timed reps per candidate, odd for an exact median (default: 7)\n");32 printf(" --seed <n> RNG seed; per-cell seeds mix it with the shape (default: 1234)\n");33 printf(" --no-cooldown do not pause/re-measure on thermal drift, only warn\n");34 printf(" --cool-drift <f> anchor drift that triggers a cooldown (default: 0.10)\n");35 printf(" --cool-eps <f> anchor tolerance to consider the GPU cool again (default: 0.03)\n");36 printf(" --cool-max-wait <s> give up cooling a cell after this many seconds (default: 120)\n");37 printf(" --cool-max-retry <n> re-measure rounds per cell before giving up (default: 2)\n");38 printf("\n");39 printf(" the table goes to stdout, all diagnostics to stderr:\n");40 printf(" %s fa-vec > rows.txt 2> sweep.log\n", argv0);41}42 43int main(int argc, char ** argv) {44 const char * tuner = nullptr;45 const char * bname = nullptr;46 tuner_opts opts;47 48 for (int i = 1; i < argc; i++) {49 const char * a = argv[i];50 if (strcmp(a, "-h") == 0 || strcmp(a, "--help") == 0) {51 usage(argv[0]);52 return 0;53 } else if (strcmp(a, "-b") == 0 && i + 1 < argc) {54 bname = argv[++i];55 } else if (strcmp(a, "--dtype") == 0 && i + 1 < argc) {56 opts.dtype_filter = argv[++i];57 } else if (strcmp(a, "--dk") == 0 && i + 1 < argc) {58 opts.dk_filter = argv[++i];59 } else if (strcmp(a, "--reps") == 0 && i + 1 < argc) {60 opts.reps = atoi(argv[++i]);61 } else if (strcmp(a, "--seed") == 0 && i + 1 < argc) {62 opts.seed = (unsigned) strtoul(argv[++i], nullptr, 10);63 } else if (strcmp(a, "--no-cooldown") == 0) {64 opts.cooldown = false;65 } else if (strcmp(a, "--cool-drift") == 0 && i + 1 < argc) {66 opts.cool_drift = atof(argv[++i]);67 } else if (strcmp(a, "--cool-eps") == 0 && i + 1 < argc) {68 opts.cool_eps = atof(argv[++i]);69 } else if (strcmp(a, "--cool-max-wait") == 0 && i + 1 < argc) {70 opts.cool_max_wait = atoi(argv[++i]);71 } else if (strcmp(a, "--cool-max-retry") == 0 && i + 1 < argc) {72 opts.cool_max_retry = atoi(argv[++i]);73 } else if (a[0] != '-' && tuner == nullptr) {74 tuner = a;75 } else {76 fprintf(stderr, "error: unrecognized or incomplete argument: %s\n\n", a);77 usage(argv[0]);78 return 1;79 }80 }81 82 if (tuner == nullptr) {83 usage(argv[0]);84 return 1;85 }86 if (opts.reps < 1) {87 fprintf(stderr, "error: --reps must be >= 1\n");88 return 1;89 }90 91 const tuner_def * t = nullptr;92 for (const auto & cand : k_tuners) {93 if (strcmp(tuner, cand.name) == 0) {94 t = &cand;95 break;96 }97 }98 if (t == nullptr) {99 fprintf(stderr, "error: unknown tuner: %s\n\n", tuner);100 usage(argv[0]);101 return 1;102 }103 104 ggml_backend_load_all();105 106 ggml_backend_dev_t dev = nullptr;107 for (size_t i = 0; i < ggml_backend_dev_count(); i++) {108 ggml_backend_dev_t d = ggml_backend_dev_get(i);109 if (bname) {110 if (strcmp(ggml_backend_dev_name(d), bname) == 0) {111 dev = d;112 break;113 }114 } else if (strncmp(ggml_backend_dev_name(d), "MTL", 3) == 0) {115 dev = d;116 break;117 }118 }119 120 if (dev == nullptr) {121 fprintf(stderr, "error: no %s device found\n", bname ? bname : "Metal");122 return 1;123 }124 125 ggml_backend_t backend = ggml_backend_dev_init(dev, nullptr);126 if (backend == nullptr) {127 fprintf(stderr, "error: failed to init backend %s\n", ggml_backend_dev_name(dev));128 return 1;129 }130 131 fprintf(stderr, "device: %s (%s)\n", ggml_backend_dev_name(dev), ggml_backend_dev_description(dev));132 133 const bool ok = t->run(backend, dev, opts);134 135 ggml_backend_free(backend);136 ggml_quantize_free();137 138 return ok ? 0 : 1;139}140 