Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
main.cpp140 linesDownload Raw Back to tuning
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