Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
fa-vec.cpp642 linesDownload Raw Back to tuning
1#include "fa-vec.h"2 3#include "bench.h"4#include "ggml-backend.h"5#include "ggml-metal-tuning.h"6#include "ggml.h"7 8#include <algorithm>9#include <cmath>10#include <cstdio>11#include <cstring>12#include <random>13#include <set>14#include <string>15#include <vector>16 17// GQA spec-decode shape: enough query heads to keep the GPU busy so the Q>1 K/V-reuse18// benefit is visible. nh KV heads, nr2 query heads each, nr3 batches.19static const int FA_NH  = 4;20static const int FA_NR2 = 8;21static const int FA_NR3 = 1;22 23struct fa_shape {24    int       dk;25    int       dv;26    int       ne01;  // query rows27    int       ne11;  // KV length28    ggml_type type_kv;29};30 31// mirrors test_flash_attn_ext::build_graph for the subset this tuner sweeps32// (mask=true, sinks=false, prec=F32, type_K==type_V, no permute)33static ggml_tensor * fa_build_graph(ggml_context * ctx, const fa_shape & s) {34    const int64_t dk_padded = GGML_PAD(s.dk, ggml_blck_size(s.type_kv));35    const int64_t dv_padded = GGML_PAD(s.dv, ggml_blck_size(s.type_kv));36 37    ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, dk_padded, s.ne01, FA_NH * FA_NR2, FA_NR3);38    ggml_set_name(q, "q");39 40    // K/V are views of a 2x-tall parent, as they are of the KV cache in production41    ggml_tensor * k0 = ggml_new_tensor_4d(ctx, s.type_kv, dk_padded, 2 * s.ne11, FA_NH, FA_NR3);42    ggml_tensor * k  = ggml_view_4d(ctx, k0, dk_padded, s.ne11, FA_NH, FA_NR3, k0->nb[1], k0->nb[2], k0->nb[3], 0);43    ggml_set_name(k, "k");44 45    ggml_tensor * v = nullptr;46    if (dk_padded == 576 && dv_padded == 512) {47        // MLA: the V cache is a sub-view of the K cache48        v = ggml_view_4d(ctx, k, dv_padded, s.ne11, FA_NH, FA_NR3, k->nb[1], k->nb[2], k->nb[3], 0);49    } else {50        ggml_tensor * v0 = ggml_new_tensor_4d(ctx, s.type_kv, dv_padded, 2 * s.ne11, FA_NH, FA_NR3);51        v                = ggml_view_4d(ctx, v0, dv_padded, s.ne11, FA_NH, FA_NR3, v0->nb[1], v0->nb[2], v0->nb[3], 0);52    }53    ggml_set_name(v, "v");54 55    ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, s.ne11, s.ne01, 1, FA_NR3);56    ggml_set_name(m, "m");57 58    ggml_tensor * out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f / sqrtf((float) s.dk), 0.0f, 0.0f);59    ggml_prec_set_acc(out, GGML_PREC_F32);60    ggml_set_name(out, "out");61 62    return out;63}64 65static uint64_t fa_op_flops(const fa_shape & s) {66    // Q*K^T is ne01 x dk x ne11, P*V is ne01 x ne11 x dv, per head67    return (uint64_t) 2 * FA_NH * FA_NR2 * s.ne01 * (s.dk + s.dv) * s.ne11 * FA_NR3;68}69 70static void fa_init_uniform(ggml_tensor * t, std::mt19937 & rng, float min, float max) {71    const size_t nels = ggml_nelements(t);72 73    std::vector<float>                    data(nels);74    std::uniform_real_distribution<float> dist(min, max);75    for (size_t i = 0; i < nels; i++) {76        data[i] = dist(rng);77    }78 79    if (t->type == GGML_TYPE_F32) {80        ggml_backend_tensor_set(t, data.data(), 0, nels * sizeof(float));81        return;82    }83 84    GGML_ASSERT(ggml_is_quantized(t->type) || t->type == GGML_TYPE_F16 || t->type == GGML_TYPE_BF16);85    GGML_ASSERT(nels % ggml_blck_size(t->type) == 0);86 87    std::vector<float> imatrix(t->ne[0], 1.0f);88    const float *      im = imatrix.data();89    if (!ggml_quantize_requires_imatrix(t->type)) {90        // when the imatrix is optional, exercise both paths; pick via one of the random numbers91        if (data[0] > 0.5f * (min + max)) {92            im = nullptr;93        }94    }95 96    const size_t blck_size = ggml_blck_size(t->type);97    const size_t n_blocks  = nels / blck_size;98 99    std::vector<uint8_t> dataq(ggml_row_size(t->type, nels));100    ggml_quantize_chunk(t->type, data.data(), dataq.data(), 0, n_blocks, blck_size, im);101 102    ggml_backend_tensor_set(t, dataq.data(), 0, dataq.size());103}104 105// mirrors init_tensor_kq_mask: f16 mask with ~20% of its blocks set to -INF or zero.106// the -INF blocks are what drives the kernel's skip-INF path, so this pattern is107// load-bearing for the timings, not just for numerics.108static void fa_init_kq_mask(ggml_tensor * t, std::mt19937 & rng, float min, float max) {109    GGML_ASSERT(t->type == GGML_TYPE_F16);110 111    const int32_t ne0 = (int32_t) t->ne[0];112    const int32_t ne1 = (int32_t) t->ne[1];113    const int32_t ne2 = (int32_t) t->ne[2];114    const int32_t ne3 = (int32_t) t->ne[3];115 116    std::vector<float>       data_f32(size_t(ne0) * ne1 * ne2 * ne3);117    std::vector<ggml_fp16_t> data_f16(size_t(ne0) * ne1 * ne2 * ne3);118 119    std::uniform_real_distribution<float> dis(min, max);120    for (size_t i = 0; i < data_f32.size(); i++) {121        data_f32[i] = dis(rng);122    }123 124    const int blck0 = 128;125    const int blck1 = 64;126 127    const int n_inf_zero_blocks = 0.2 * (ne0 * ne1 * ne2 * ne3) / (blck0 * blck1);128 129    for (int b = 0; b < n_inf_zero_blocks; b++) {130        const int p3 = (int) (rng() % ne3);131        const int p2 = (int) (rng() % ne2);132        const int p1 = (int) (rng() % ne1);133        const int p0 = (int) (rng() % ne0);134 135        const bool inf = rng() & 1;136 137        for (int i1 = 0; i1 < blck1 && p1 + i1 < ne1; i1++) {138            const int idx = p3 * ne2 * ne1 * ne0 + p2 * ne1 * ne0 + (p1 + i1) * ne0 + p0;139 140            for (int i0 = 0; i0 < blck0 && p0 + i0 < ne0; i0++) {141                data_f32[idx + i0] = inf ? -INFINITY : 0.0f;142            }143        }144    }145 146    ggml_fp32_to_fp16_row(data_f32.data(), data_f16.data(), ne0 * ne1 * ne2 * ne3);147 148    ggml_backend_tensor_set(t, data_f16.data(), 0, data_f16.size() * sizeof(ggml_fp16_t));149}150 151static unsigned fa_cell_seed(const fa_shape & s, unsigned base) {152    unsigned h = base;153    for (int v : { s.dk, s.dv, s.ne01, s.ne11, (int) s.type_kv }) {154        h = h * 1000003u + (unsigned) v;155    }156    return h;157}158 159static void fa_init_tensors(ggml_context * ctx, const fa_shape & s, unsigned base_seed) {160    std::mt19937 rng(fa_cell_seed(s, base_seed));161 162    for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {163        if (t->view_src != NULL) {164            continue;  // views share their parent's data165        }166        if (strcmp(t->name, "m") == 0) {167            fa_init_kq_mask(t, rng, -1.0f, 1.0f);168        } else {169            fa_init_uniform(t, rng, -1.0f, 1.0f);170        }171    }172}173 174using set_override_t   = void (*)(int, int);175using clear_override_t = void (*)(void);176using bucket_t         = int (*)(int64_t);177using baseline_ne_t    = int (*)(int, int);178using device_token_t   = const char * (*) (ggml_backend_dev_t);179 180struct fa_procs {181    set_override_t   set_ov      = nullptr;182    clear_override_t clr_ov      = nullptr;183    bucket_t         ne11_bucket = nullptr;184    bucket_t         ne01_bucket = nullptr;185    baseline_ne_t    baseline_ne = nullptr;186    device_token_t   dev_token   = nullptr;187 188    bool ok() const { return set_ov && clr_ov && ne11_bucket && ne01_bucket && baseline_ne && dev_token; }189};190 191static fa_procs fa_resolve_procs(ggml_backend_dev_t dev) {192    ggml_backend_reg_t reg = ggml_backend_dev_backend_reg(dev);193 194    fa_procs p;195    p.set_ov = (set_override_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_set_fa_vec_override");196    p.clr_ov =197        (clear_override_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_clear_fa_vec_override");198    p.ne11_bucket = (bucket_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_fa_vec_ne11_bucket");199    p.ne01_bucket = (bucket_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_fa_vec_ne01_bucket");200    p.baseline_ne =201        (baseline_ne_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_fa_vec_baseline_ne");202    p.dev_token = (device_token_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_device_token");203 204    return p;205}206 207static bool fa_filter_has(const char * filter, const char * name) {208    if (!filter) {209        return true;210    }211 212    const std::string f = std::string(",") + filter + ",";213 214    return f.find(std::string(",") + name + ",") != std::string::npos;215}216 217struct fa_cand {218    int Q, NE;219};220 221struct fa_point {222    int                 dk, dv, ne11, ne01;223    std::vector<double> t;224};225 226// base_i identifies the (Q=1, baseline NE) anchor configuration.227static std::vector<fa_cand> fa_build_cands(const fa_procs & procs, int dk, int dv, int & base_i) {228    const int base_ne = procs.baseline_ne(dk, dv);229 230    std::vector<fa_cand> cands;231    base_i = -1;232    for (int ne : ggml_metal_tuning::fa_vec_legal_ne(dk, dv)) {233        for (int Q : { 1, 2, 4 }) {234            if (Q == 1 && ne == base_ne) {235                base_i = (int) cands.size();236            }237            cands.push_back({ Q, ne });238        }239    }240    GGML_ASSERT(base_i >= 0);241 242    return cands;243}244 245bool tuner_fa_vec_run(ggml_backend_t backend, ggml_backend_dev_t dev, const tuner_opts & opts) {246    const fa_procs procs = fa_resolve_procs(dev);247    if (!procs.ok()) {248        fprintf(stderr, "error: metal fa_vec tuning procs unavailable\n");249        return false;250    }251 252    const char * dev_token = procs.dev_token(dev);253 254    struct shape_t {255        int dk, dv;256    };257 258    const shape_t shapes[] = {259        { 32,  32  },260        { 64,  64  },261        { 96,  96  },262        { 128, 128 },263        { 192, 192 },264        { 192, 128 },265        { 256, 256 },266        { 320, 256 },267        { 512, 512 },268        { 576, 512 }269    };270    // nsg is a pipeline specialization constant (1 up to ne11=2048, 2 up to 4096, 4 above), so ne11271    // bucket 1 takes two samples to cover both of its regimes. Bucket 0 is not sampled at all: the272    // runtime leaves short KV at baseline, so no measurement there can reach the table.273    const int ne11_rep[] = { 2048, 3072, 8192, 32768 };274    const int ne01_rep[] = { 1, 2, 3, 4, 5, 6, 7, 8, 16 };  // point buckets (1-4) + tail mod-4 cycle + anchor275 276    struct dtype_t {277        ggml_type    type;278        const char * token;279    };280 281    const dtype_t dtypes[] = {282        { GGML_TYPE_F16,  "GGML_TYPE_F16"  },283        { GGML_TYPE_Q4_0, "GGML_TYPE_Q4_0" },284        { GGML_TYPE_Q4_1, "GGML_TYPE_Q4_1" },285        { GGML_TYPE_Q5_0, "GGML_TYPE_Q5_0" },286        { GGML_TYPE_Q5_1, "GGML_TYPE_Q5_1" },287        { GGML_TYPE_Q8_0, "GGML_TYPE_Q8_0" },288    };289 290    const double TUNE_TAU   = 0.05;  // max POINTWISE regret to ride a domain default291    const double TUNE_THETA = 1.05;  // min AGGREGATE bucket speedup vs baseline to tune at all292 293    const cooldown_opts cool = {294        opts.cooldown, opts.cool_drift, opts.cool_eps, opts.cool_max_wait, opts.cool_max_retry,295    };296 297    fprintf(stderr, "seed=%u reps=%d cooldown=%s (drift=%.2f eps=%.2f max_wait=%ds max_retry=%d)\n", opts.seed,298            opts.reps, cool.enabled ? "on" : "off", cool.drift, cool.eps, cool.max_wait, cool.max_retry);299    fprintf(stderr, "device token: %s\n", dev_token);300 301    int n_untrusted = 0;302 303    // stdout carries nothing but table rows, so the whole stream pastes into fa_vec_tuned_table304    for (const auto & dtype : dtypes) {305        const ggml_type type_kv = dtype.type;306        if (!fa_filter_has(opts.dtype_filter, ggml_type_name(type_kv))) {307            continue;308        }309 310        fprintf(stderr, "\n### dtype=%s\n", ggml_type_name(type_kv));311 312        std::vector<fa_point> pts;313 314        for (auto s : shapes) {315            if (!fa_filter_has(opts.dk_filter, std::to_string(s.dk).c_str())) {316                continue;317            }318 319            int                  base_i = 0;320            std::vector<fa_cand> cands  = fa_build_cands(procs, s.dk, s.dv, base_i);321 322            for (int ne11 : ne11_rep) {323                for (int ne01 : ne01_rep) {324                    const fa_shape sh = { s.dk, s.dv, ne01, ne11, type_kv };325 326                    perf_cell cell = build_perf_cell(327                        backend, [&](ggml_context * ctx) { return fa_build_graph(ctx, sh); },328                        [&](ggml_context * ctx) { fa_init_tensors(ctx, sh, opts.seed); },329                        [&](ggml_tensor *) { return fa_op_flops(sh); });330 331                    if (cell.gf == nullptr) {332                        continue;333                    }334 335                    // randomize candidate order to decorrelate thermal drift across the cell336                    std::vector<int> order((size_t) cands.size());337                    for (size_t i = 0; i < order.size(); ++i) {338                        order[i] = (int) i;339                    }340                    std::shuffle(order.begin(), order.end(), std::mt19937(fa_cell_seed(sh, opts.seed)));341 342                    char label[128];343                    snprintf(label, sizeof(label), "dk=%d ne11=%d", s.dk, ne11);344 345                    cell_result r = measure_cell(346                        backend, cell, opts.reps, order, [&](int i) { procs.set_ov(cands[i].Q, cands[i].NE); },347                        [&]() { procs.clr_ov(); }, base_i, cool, label);348 349                    if (r.anchor_min > 0.0) {350                        fprintf(stderr, "# noise dk=%d dv=%d ne11=%d ne01=%d spread=%.1f%%\n", s.dk, s.dv, ne11, ne01,351                                100.0 * (r.anchor_max - r.anchor_min) / r.anchor_min);352                    }353 354                    if (!r.trusted) {355                        n_untrusted++;356                        fprintf(stderr, "# DROP untrusted cell dk=%d dv=%d ne11=%d ne01=%d\n", s.dk, s.dv, ne11, ne01);357                        continue;358                    }359 360                    int best_i = -1;361                    for (size_t i = 0; i < cands.size(); ++i) {362                        if (r.t[i] > 0.0 && (best_i < 0 || r.t[i] < r.t[best_i])) {363                            best_i = (int) i;364                        }365                    }366                    const double base_t = r.t[base_i];367                    const bool   keep   = best_i >= 0 && base_t > 0.0 && r.t[best_i] < base_t * 0.98;368 369                    fprintf(stderr, "# dtype=%s dk=%d dv=%d ne11=%d ne01=%d:", ggml_type_name(type_kv), s.dk, s.dv,370                            ne11, ne01);371                    for (size_t i = 0; i < cands.size(); ++i) {372                        fprintf(stderr, "  Q%dNE%d=%.1f%s", cands[i].Q, cands[i].NE, r.t[i],373                                (int) i == best_i ? "*" : "");374                    }375                    if (keep) {376                        fprintf(stderr, "  => Q%d,NE%d  %.2fx\n", cands[best_i].Q, cands[best_i].NE,377                                base_t / r.t[best_i]);378                    } else {379                        fprintf(stderr, "  => baseline\n");380                    }381 382                    pts.push_back({ s.dk, s.dv, ne11, ne01, r.t });383                }384            }385        }386 387        // compress into pasteable rows. per (dk,dv) and ne01 domain {decode==1, batch>=2},388        // emit one ne11-collapsed default cfg (ne11_b=-1) plus a per-bucket exception wherever the389        // default's pointwise regret vs the bucket target exceeds TUNE_TAU, or the default is not390        // admissible for that bucket (see never_slower / admissible below).391        std::vector<std::string> rows_out;392        char                     rbuf[192];393 394        for (auto s : shapes) {395            if (!fa_filter_has(opts.dk_filter, std::to_string(s.dk).c_str())) {396                continue;397            }398 399            int                  base_i = 0;400            std::vector<fa_cand> cands  = fa_build_cands(procs, s.dk, s.dv, base_i);401 402            struct bkt_t {403                int                           b11, b01, Ti;404                std::vector<double>           agg;405                std::vector<const fa_point *> bp;406            };407 408            // A config may represent a bucket only if it is no slower than baseline at every point that409            // bucket covers. The aggregate gate below sums absolute times, so it can pass on the aligned410            // and deep points while a misaligned ne01 pays the mod-Q padding. Nothing measured, nothing411            // proven: a bucket with no surviving sample admits baseline only.412            auto never_slower = [&](const std::vector<const fa_point *> & bp, int i) {413                if (i == base_i) {414                    return true;415                }416                if (bp.empty()) {417                    return false;418                }419                for (const auto * p : bp) {420                    if (p->t[i] <= 0.0 || p->t[base_i] <= 0.0 || p->t[i] > p->t[base_i]) {421                        return false;422                    }423                }424                return true;425            };426 427            // The padded-row waste ceil(n/Q)*Q/n is largest at the smallest ne01 of each residue class428            // mod Q, so one of a bucket's first Q values carries the worst padding it can ever see, and429            // that value has to be sampled. Otherwise the bucket bounds nothing: a config picked on the430            // aligned ne01=8,16 says nothing about ne01=9. This covers the padding term only - the431            // per-row cost varies with ne01 too - so it is a floor on the evidence, not a proof.432            auto admissible = [&](const std::vector<const fa_point *> & bp, int b01, int i) {433                if (!never_slower(bp, i)) {434                    return false;435                }436                const int Q = cands[i].Q;437                if (Q == 1) {438                    return true;  // one row per threadgroup, no padding to witness439                }440                int lo = bp[0]->ne01;441                for (const auto * p : bp) {442                    lo = std::min(lo, p->ne01);443                }444                while (lo > 1 && procs.ne01_bucket(lo - 1) == b01) {445                    lo--;  // walk down to where this bucket's runtime domain starts446                }447                int    wit  = lo;448                double wmax = 0.0;449                for (int n = lo; n < lo + Q && procs.ne01_bucket(n) == b01; ++n) {450                    const int    padded = ((n + Q - 1) / Q) * Q;451                    const double w      = (double) padded / n;452                    if (w > wmax) {453                        wmax = w;454                        wit  = n;455                    }456                }457                for (const auto * p : bp) {458                    if (p->ne01 == wit) {459                        return true;460                    }461                }462                return false;463            };464 465            std::set<std::pair<int, int>> buckets;466            for (int ne11 : ne11_rep) {467                const int b11 = procs.ne11_bucket(ne11);468                if (b11 == 0) {469                    continue;470                }471                for (int ne01 : ne01_rep) {472                    buckets.insert({ b11, procs.ne01_bucket(ne01) });473                }474            }475 476            std::vector<bkt_t> bks;477            for (const auto & bb : buckets) {478                const int b11 = bb.first, b01 = bb.second;479 480                std::vector<const fa_point *> bp;481                for (const auto & p : pts) {482                    if (p.dk == s.dk && p.dv == s.dv && procs.ne11_bucket(p.ne11) == b11 &&483                        procs.ne01_bucket(p.ne01) == b01) {484                        bp.push_back(&p);485                    }486                }487 488                fprintf(stderr, "# bucket dk=%d dv=%d ne11_b=%d ne01_b=%d samples=%zu\n", s.dk, s.dv, b11, b01,489                        bp.size());490                if (bp.empty()) {491                    // nothing to check a config against, so pin the bucket to baseline instead of492                    // letting the ne11-collapsed domain default ride in unmeasured493                    fprintf(stderr, "# WARN empty bucket dk=%d dv=%d ne11_b=%d ne01_b=%d -> baseline\n", s.dk, s.dv,494                            b11, b01);495                    bks.push_back({ b11, b01, base_i, std::vector<double>(cands.size(), 0.0), {} });496                    continue;497                }498 499                std::vector<double> agg(cands.size(), 0.0), worst(cands.size(), 0.0);500                for (const auto * p : bp) {501                    double bestt = 0.0;502                    for (size_t i = 0; i < cands.size(); ++i) {503                        if (p->t[i] > 0.0 && (bestt == 0.0 || p->t[i] < bestt)) {504                            bestt = p->t[i];505                        }506                    }507                    for (size_t i = 0; i < cands.size(); ++i) {508                        agg[i] += p->t[i];509                        if (p->t[i] > 0.0 && bestt > 0.0) {510                            worst[i] = std::max(worst[i], p->t[i] / bestt);511                        }512                    }513                }514 515                int robust = -1, oracle_pick = -1;516                for (size_t i = 0; i < cands.size(); ++i) {517                    auto tighter = [&](int j) {518                        return j < 0 || worst[i] < worst[j] ||519                               (worst[i] == worst[j] && (cands[i].Q < cands[j].Q ||520                                                         (cands[i].Q == cands[j].Q && cands[i].NE < cands[j].NE)));521                    };522                    if (tighter(oracle_pick)) {523                        oracle_pick = (int) i;524                    }525                    if (admissible(bp, b01, (int) i) && tighter(robust)) {526                        robust = (int) i;527                    }528                }529 530                const bool tune = robust != base_i && agg[base_i] > 0.0 && agg[robust] > 0.0 &&531                                  agg[base_i] / agg[robust] >= TUNE_THETA;532 533                // report what the no-harm rule cost this bucket, but only when it changed the outcome:534                // a sweep on another machine then shows where the winner loses, instead of just535                // emitting a smaller table536                const bool refused = oracle_pick != robust && oracle_pick != base_i && agg[base_i] > 0.0 &&537                                     agg[oracle_pick] > 0.0 && agg[base_i] / agg[oracle_pick] >= TUNE_THETA;538                if (refused) {539                    double over = 0.0;540                    int    at11 = 0, at01 = 0;541                    for (const auto * p : bp) {542                        if (p->t[base_i] > 0.0 && p->t[oracle_pick] / p->t[base_i] - 1.0 > over) {543                            over = p->t[oracle_pick] / p->t[base_i] - 1.0;544                            at11 = p->ne11;545                            at01 = p->ne01;546                        }547                    }548                    if (over > 0.0) {549                        fprintf(stderr,550                                "# reject dk=%d dv=%d ne11_b=%d ne01_b=%d Q%dNE%d: +%.2f%% vs baseline at "551                                "ne11=%d ne01=%d\n",552                                s.dk, s.dv, b11, b01, cands[oracle_pick].Q, cands[oracle_pick].NE, 100.0 * over, at11,553                                at01);554                    } else {555                        fprintf(stderr, "# reject dk=%d dv=%d ne11_b=%d ne01_b=%d Q%dNE%d: no padding witness\n", s.dk,556                                s.dv, b11, b01, cands[oracle_pick].Q, cands[oracle_pick].NE);557                    }558                }559 560                bks.push_back({ b11, b01, tune ? robust : base_i, agg, bp });561            }562 563            // pointwise regret of default cfg d vs the bucket target: a ratio-of-sums lets a564            // default that wins on aligned ne01 hide a large penalty on a misaligned point565            auto reg_pointwise = [&](const bkt_t * b, int d) {566                double r = 0.0;567                for (const auto * p : b->bp) {568                    const double td = p->t[d], tT = p->t[b->Ti];569                    if (td > 0.0 && tT > 0.0) {570                        r = std::max(r, td / tT - 1.0);571                    }572                }573                return r;574            };575 576            for (int dom = 0; dom <= 1; ++dom) {  // 0 = decode (ne01==1), 1 = batch (ne01>=2)577                std::vector<const bkt_t *> db;578                for (const auto & b : bks) {579                    if ((dom == 0) == (b.b01 == 0)) {580                        db.push_back(&b);581                    }582                }583                if (db.empty()) {584                    continue;585                }586 587                // default cfg = the one minimizing (#rows, total achieved time, Q, NE)588                int    bestD = -1, bestRows = 1 << 30;589                double bestTot = 0.0;590                for (size_t d = 0; d < cands.size(); ++d) {591                    int    rows = ((int) d != base_i) ? 1 : 0;592                    double tot  = 0.0;593                    for (const auto * b : db) {594                        if (reg_pointwise(b, (int) d) > TUNE_TAU || !admissible(b->bp, b->b01, (int) d)) {595                            rows++;596                            tot += b->agg[b->Ti];597                        } else {598                            tot += b->agg[d];599                        }600                    }601                    const bool better =602                        bestD < 0 || rows < bestRows ||603                        (rows == bestRows &&604                         (tot < bestTot ||605                          (tot == bestTot && (cands[d].Q < cands[bestD].Q ||606                                              (cands[d].Q == cands[bestD].Q && cands[d].NE < cands[bestD].NE)))));607                    if (better) {608                        bestD    = (int) d;609                        bestRows = rows;610                        bestTot  = tot;611                    }612                }613 614                if (bestD != base_i) {615                    snprintf(rbuf, sizeof(rbuf), "    { { %s, %s, %d, %d, -1, %d }, { %d, %d } },", dev_token,616                             dtype.token, s.dk, s.dv, dom, cands[bestD].Q, cands[bestD].NE);617                    rows_out.emplace_back(rbuf);618                }619                for (const auto * b : db) {620                    if (reg_pointwise(b, bestD) <= TUNE_TAU && admissible(b->bp, b->b01, bestD)) {621                        continue;622                    }623                    snprintf(rbuf, sizeof(rbuf), "    { { %s, %s, %d, %d, %d, %d }, { %d, %d } },", dev_token,624                             dtype.token, s.dk, s.dv, b->b11, b->b01, cands[b->Ti].Q, cands[b->Ti].NE);625                    rows_out.emplace_back(rbuf);626                }627            }628        }629 630        for (const auto & r : rows_out) {631            printf("%s\n", r.c_str());632        }633        fflush(stdout);634    }635 636    if (n_untrusted > 0) {637        fprintf(stderr, "\n%d cells excluded as untrusted (see DROP lines above)\n", n_untrusted);638    }639 640    return true;641}642