Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
test-backend-sampler.cpp1167 linesDownload Raw Back to tests
1#include "ggml.h"2#include "llama.h"3#include "llama-cpp.h"4#include "get-model.h"5#include "common.h"6 7#ifdef NDEBUG8#undef NDEBUG9#endif10 11#include <algorithm>12#include <cstdlib>13#include <cstring>14#include <fstream>15#include <map>16#include <string>17#include <unordered_map>18#include <vector>19 20struct test_args {21    std::string model;22    std::string test;23    std::string device = "auto";24};25 26struct test_params {27    llama_model_ptr model;28};29 30static llama_model_ptr load_model(const test_args & args) {31    auto mparams = llama_model_default_params();32 33    ggml_backend_dev_t devs[2] = { nullptr, nullptr };34 35    if (args.device != "auto") {36        if (args.device == "gpu") {37            devs[0] = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_GPU);38 39            if (devs[0] == nullptr) {40                fprintf(stderr, "Error: GPU requested but not available\n");41                return nullptr;42            }43 44            mparams.n_gpu_layers = 999;45        } else if (args.device == "cpu") {46            devs[0] = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU);47 48            mparams.n_gpu_layers = 0;49        } else {50            fprintf(stderr, "Error: invalid device '%s'\n", args.device.c_str());51            return nullptr;52        }53 54        mparams.devices = devs;55 56        fprintf(stderr, "Using device: %s\n", ggml_backend_dev_name(devs[0]));57    }58 59    llama_model_ptr res;60 61    res.reset(llama_model_load_from_file(args.model.c_str(), mparams));62 63    if (!res) {64        fprintf(stderr, "Warning: failed to load model '%s', skipping test\n", args.model.c_str());65        return nullptr;66    }67 68    return res;69}70 71struct test_context {72    llama_context_ptr ctx;73 74    int n_vocab = 0;75 76    const llama_vocab * vocab = nullptr;77 78    std::unordered_map<llama_seq_id, int32_t> seq_positions;79    std::unordered_map<llama_seq_id, int32_t> last_batch_info;80 81    test_context(const test_params & params, std::vector<llama_sampler_seq_config> & configs, int32_t n_seq_max = -1) {82        auto * model = params.model.get();83 84        GGML_ASSERT(model);85        GGML_ASSERT(!ctx);86 87        llama_context_params cparams = llama_context_default_params();88        cparams.n_ctx = 512;89        cparams.n_batch = 512;90        cparams.samplers = configs.data();91        cparams.n_samplers = configs.size();92        cparams.kv_unified = true;93 94        // If n_seq_max is not specified, calculate it from configs95        if (n_seq_max < 0) {96            int32_t max_seq_id = 0;97            for (const auto & config : configs) {98                max_seq_id = std::max(config.seq_id, max_seq_id);99            }100            cparams.n_seq_max = max_seq_id + 1;101        } else {102            cparams.n_seq_max = n_seq_max;103        }104 105        ctx.reset(llama_init_from_model(model, cparams));106        if (!ctx) {107            throw std::runtime_error("failed to create context");108        }109 110        llama_set_warmup(ctx.get(), false);111 112        vocab = llama_model_get_vocab(model);113        n_vocab = llama_vocab_n_tokens(vocab);114    }115 116    bool decode(const std::map<llama_seq_id, std::string> & prompts) {117        GGML_ASSERT(ctx);118 119        last_batch_info.clear();120        llama_batch batch = llama_batch_init(512, 0, prompts.size());121 122        for (const auto & [seq_id, prompt] : prompts) {123            std::vector<llama_token> tokens;124            tokens.push_back(llama_vocab_bos(vocab));125 126            std::vector<llama_token> prompt_tokens(32);127            int n_tokens = llama_tokenize(vocab, prompt.c_str(), prompt.length(),128                                           prompt_tokens.data(), prompt_tokens.size(),129                                           false, false);130            if (n_tokens < 0) {131                fprintf(stderr, "Warning: tokenization failed for seq_id %d\n", seq_id);132                llama_batch_free(batch);133                return false;134            }135 136            for (int i = 0; i < n_tokens; i++) {137                tokens.push_back(prompt_tokens[i]);138            }139 140            if (seq_positions.find(seq_id) == seq_positions.end()) {141                seq_positions[seq_id] = 0;142            }143 144            int32_t start_pos = seq_positions[seq_id];145            for (size_t i = 0; i < tokens.size(); i++) {146                common_batch_add(batch, tokens[i], start_pos + i, { seq_id }, i == tokens.size() - 1);147            }148 149            seq_positions[seq_id] = start_pos + tokens.size();150        }151 152 153        printf("Batch contents:\n");154        printf("n_tokens: %d\n", batch.n_tokens);155        for (int i = 0; i < batch.n_tokens; i++) {156            printf("token[%d]: tok=%-5d, pos=%d, n_seq_id=%d, seq_ids=[", i, batch.token[i], batch.pos[i], batch.n_seq_id[i]);157 158            for (int j = 0; j < batch.n_seq_id[i]; j++) {159                printf("%d%s", batch.seq_id[i][j], j < batch.n_seq_id[i]-1 ? ", " : "");160            }161            printf("], logits=%d\n", batch.logits[i]);162        }163 164        if (llama_decode(ctx.get(), batch) != 0) {165            fprintf(stderr, "Warning: llama_decode failed\n");166            llama_batch_free(batch);167            return false;168        }169 170        // Build mapping from seq id to batch token idx171        for (int i = 0; i < batch.n_tokens; i++) {172            if (batch.logits[i]) {173                llama_seq_id seq_id = batch.seq_id[i][0];174                last_batch_info[seq_id] = i;175            }176        }177 178        llama_batch_free(batch);179        return true;180    }181 182    int32_t idx_for_seq(llama_seq_id seq_id) {183        auto it = last_batch_info.find(seq_id);184        if (it == last_batch_info.end()) {185            fprintf(stderr, "Error: no batch index found for seq_id %d\n", seq_id);186            return -1;187        }188        return it->second;189    }190 191    void update_batch_info(const llama_batch & batch) {192        last_batch_info.clear();193        for (int i = 0; i < batch.n_tokens; i++) {194            if (batch.logits[i]) {195                llama_seq_id cur_seq = batch.seq_id[i][0];196                last_batch_info[cur_seq] = i;197            }198        }199    }200 201    bool decode_token(llama_token token, llama_seq_id seq_id = 0) {202        GGML_ASSERT(ctx);203 204        llama_batch batch = llama_batch_init(1, 0, 1);205        int32_t pos = seq_positions[seq_id];206        common_batch_add(batch, token, pos, { seq_id }, true);207 208        if (llama_decode(ctx.get(), batch) != 0) {209            fprintf(stderr, "Warning: llama_decode failed for token %d in seq %d\n", token, seq_id);210            llama_batch_free(batch);211            return false;212        }213 214        update_batch_info(batch);215 216        seq_positions[seq_id]++;217        llama_batch_free(batch);218 219        return true;220    }221 222    bool decode_tokens(const std::map<llama_seq_id, llama_token> & seq_tokens) {223        GGML_ASSERT(ctx);224 225        llama_batch batch = llama_batch_init(seq_tokens.size(), 0, seq_tokens.size());226 227        for (const auto & [seq_id, token] : seq_tokens) {228            int32_t pos = seq_positions[seq_id];229            common_batch_add(batch, token, pos, { seq_id }, true);230        }231 232        if (llama_decode(ctx.get(), batch) != 0) {233            fprintf(stderr, "Warning: llama_decode failed for batch tokens\n");234            llama_batch_free(batch);235            return false;236        }237 238        for (const auto & [seq_id, _] : seq_tokens) {239            seq_positions[seq_id]++;240        }241 242        update_batch_info(batch);243 244        llama_batch_free(batch);245 246        return true;247    }248 249    std::string token_to_piece(llama_token token, bool special) const {250        std::string piece;251        piece.resize(piece.capacity());  // using string internal cache, 15 bytes + '\n'252        const int n_chars = llama_token_to_piece(vocab, token, &piece[0], piece.size(), 0, special);253        if (n_chars < 0) {254            piece.resize(-n_chars);255            int check = llama_token_to_piece(vocab, token, &piece[0], piece.size(), 0, special);256            GGML_ASSERT(check == -n_chars);257        } else {258            piece.resize(n_chars);259        }260 261        return piece;262    }263};264 265static void test_backend_greedy_sampling(const test_params & params) {266    const int seq_id = 0;267 268    struct llama_sampler_chain_params backend_sampler_params = llama_sampler_chain_default_params();269    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_sampler_params));270 271    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_greedy());272    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};273 274    test_context test_ctx(params, backend_sampler_configs);275 276    if (!test_ctx.decode({{seq_id, "Some"}})) {277        GGML_ASSERT(false && "Failed to decode token");278    }279 280    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);281 282    llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);283    printf("greedy sampled id:%d, string:'%s'\n", token, test_ctx.token_to_piece(token, false).c_str());284    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);285 286    token = llama_get_sampled_token_ith(test_ctx.ctx.get(), -1);287    printf("greedy sampled id:%d, string:'%s'\n", token, test_ctx.token_to_piece(token, false).c_str());288    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);289 290    for (int i = 0; i < 10; i++) {291        int32_t loop_idx = test_ctx.idx_for_seq(seq_id);292        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), loop_idx);293        printf("Generation step %d: token id:%d, string: %s\n", i, token, test_ctx.token_to_piece(token, false).c_str());294        if (!test_ctx.decode_token(token, 0)) {295            GGML_ASSERT(false && "Failed to decode token");296        }297    }298}299 300static void test_backend_top_k_sampling(const test_params & params) {301    const int seq_id = 0;302    const int32_t k = 8;303    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();304    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));305    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_top_k(k));306    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};307 308    test_context test_ctx(params, backend_sampler_configs);309 310    if (!test_ctx.decode({{seq_id, "Hello"}})) {311        GGML_ASSERT(false && "Failed to decode token");312    }313 314    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);315 316    float * logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx);317    uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);318    for (size_t i = 0; i < n_logits; ++i) {319        printf("top_k logit[%zu] = %.6f\n", i, logits[i]);320    }321 322    llama_token * candidates = llama_get_sampled_candidates_ith(test_ctx.ctx.get(), batch_idx);323    uint32_t n_candidates = llama_get_sampled_candidates_count_ith(test_ctx.ctx.get(), batch_idx);324    for (size_t i = 0; i < n_candidates; ++i) {325        printf("top_k candidate[%zu] = %d : %s\n", i, candidates[i],326               test_ctx.token_to_piece(candidates[i], false).c_str());327    }328 329    // Sample using CPU sampler for verification that it is possible to do hybrid330    // sampling, first top_k on the backend and then dist on the CPU.331    struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();332    llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));333    GGML_ASSERT(chain->iface->backend_apply != nullptr);334 335    llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(18));336    llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);337    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);338 339    printf("backend top-k hybrid sampling test PASSED\n");340}341 342static void test_backend_temp_sampling(const test_params & params) {343    {344        const float temp_0 = 0.8f;345        struct llama_sampler_chain_params backend_chain_params_0 = llama_sampler_chain_default_params();346        llama_sampler_ptr backend_sampler_chain_0(llama_sampler_chain_init(backend_chain_params_0));347        llama_sampler_chain_add(backend_sampler_chain_0.get(), llama_sampler_init_temp(temp_0));348 349        const float temp_1 = 0.1f;350        struct llama_sampler_chain_params backend_chain_params_1 = llama_sampler_chain_default_params();351        llama_sampler_ptr backend_sampler_chain_1(llama_sampler_chain_init(backend_chain_params_1));352        llama_sampler_chain_add(backend_sampler_chain_1.get(), llama_sampler_init_temp(temp_1));353 354        std::vector<llama_sampler_seq_config> backend_sampler_configs = {355            { 0, backend_sampler_chain_0.get() },356            { 1, backend_sampler_chain_1.get() }357        };358 359        test_context test_ctx(params, backend_sampler_configs);360 361        if (!test_ctx.decode({{0, "Some where over the"}, {1, "Once upon a"}})) {362            GGML_ASSERT(false && "Failed to decode token");363        }364 365        // Verify sequence 0366        {367            int32_t batch_idx = test_ctx.idx_for_seq(0);368            int n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);369            GGML_ASSERT(n_logits == test_ctx.n_vocab);370 371            // Sample from sequence 0 using CPU sampler372            struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();373            llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));374            llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(18));375 376            llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);377            const std::string token_str = test_ctx.token_to_piece(token, false);378            printf("Sequence 0 sampled token id:%d, string: '%s'\n", token, token_str.c_str());379            GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);380        }381 382 383        // Verify sequence 1384        {385            int32_t batch_idx = test_ctx.idx_for_seq(1);386 387            // Sample from sequence 1 using CPU sampler388            struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();389            llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));390            llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(18));391 392            llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);393            const std::string token_str = test_ctx.token_to_piece(token, false);394            printf("Sequence 1 sampled token id:%d, string: '%s'\n", token, token_str.c_str());395            GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);396        }397    }398 399    // lambda for testing non-positive temperature values.400    auto test_argmax_temp = [&](float temp) {401        printf("\nTesting temperature = %.1f\n", temp);402 403        int seq_id = 0;404        struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();405        llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));406        llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_temp(temp));407 408        std::vector<llama_sampler_seq_config> backend_sampler_configs = {409            { seq_id, backend_sampler_chain.get() },410        };411 412        test_context test_ctx(params, backend_sampler_configs);413 414        if (!test_ctx.decode({{seq_id, "Once"}})) {415            GGML_ASSERT(false && "Failed to decode token");416        }417 418        int32_t batch_idx = test_ctx.idx_for_seq(seq_id);419 420        uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);421        GGML_ASSERT(n_logits == 1);422    };423 424    test_argmax_temp(0.0f);425    test_argmax_temp(-1.0f);426 427    printf("backend temp sampling test PASSED\n");428}429 430static void test_backend_temp_ext_sampling(const test_params & params) {431    {432        int seq_id = 0;433        const float temp = 0.8f;434        const float delta = 0.5f;435        const float exponent = 1.5f;436        struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();437        llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));438        llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_temp_ext(temp, delta, exponent));439 440        std::vector<llama_sampler_seq_config> backend_sampler_configs = {441            { seq_id, backend_sampler_chain.get() },442        };443 444        test_context test_ctx(params, backend_sampler_configs);445 446        if (!test_ctx.decode({{seq_id, "Once upon a"}})) {447            GGML_ASSERT(false && "Failed to decode token");448        }449 450        // Verify sequence 0451        {452            int32_t batch_idx = test_ctx.idx_for_seq(seq_id);453            int n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);454            GGML_ASSERT(n_logits == test_ctx.n_vocab);455        }456    }457 458    // lambda for testing non-positive temp/delta/exponent values.459    auto test_argmax_temp = [&](float temp, float delta, float exponent) {460        printf("\nTesting temperature = %.1f, delta = %1.f, exponent = %1.f\n", temp, delta, exponent);461 462        int seq_id = 0;463        struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();464        llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));465        llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_temp_ext(temp, delta, exponent));466 467        std::vector<llama_sampler_seq_config> backend_sampler_configs = {468            { seq_id, backend_sampler_chain.get() },469        };470 471        test_context test_ctx(params, backend_sampler_configs);472 473        if (!test_ctx.decode({{seq_id, "Once"}})) {474            GGML_ASSERT(false && "Failed to decode token");475        }476 477        int32_t batch_idx = test_ctx.idx_for_seq(seq_id);478 479        uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);480 481        if (temp <= 0.0f && delta >= 0.0f) {482            GGML_ASSERT(n_logits == 1);483        } else {484            GGML_ASSERT(n_logits == (uint32_t) test_ctx.n_vocab);485        }486    };487 488    test_argmax_temp(0.0f,  0.3f, 1.0f); // Greedy (temp=0)489    test_argmax_temp(-1.0f, 0.3f, 2.0f); // Greedy (temp<0)490    test_argmax_temp(0.8f,  0.0f, 2.0f); // Temperature scaling491 492    printf("backend temp_ext sampling test PASSED\n");493}494 495static void test_backend_min_p_sampling(const test_params & params) {496    const int seq_id = 0;497    const float p = 0.1;498    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();499    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));500    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_min_p(p, 0));501    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};502 503    test_context test_ctx(params, backend_sampler_configs);504 505    if (!test_ctx.decode({{seq_id, "Hello"}})) {506        GGML_ASSERT(false && "Failed to decode token");507    }508 509    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);510 511    float * logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx);512    uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);513 514    // Print the logits that are above the min-p threshold515    std::vector<float> filtered_logits;516    for (size_t i = 0; i < n_logits; ++i) {517        if (logits[i] > -1e9f) {518            filtered_logits.push_back(logits[i]);519            //printf("min_p logit[%zu] = %.6f\n", i, logits[i]);520        }521    }522    GGML_ASSERT(filtered_logits.size() < (size_t) test_ctx.n_vocab);523 524    // Sample using CPU sampler for verification to inspect they are reasonable525    struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();526    llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));527    llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(88));528 529    llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);530    const std::string token_str = test_ctx.token_to_piece(token, false);531    printf("min-p cpu sampled token id:%d, string: '%s'\n", token, token_str.c_str());532    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);533 534    // Decode and sample 10 more tokens535    for (int i = 0; i < 10; i++) {536        int32_t loop_idx = test_ctx.idx_for_seq(seq_id);537        llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), loop_idx);538        printf("min-p gen step %d: token id :%5.d, string: %s\n", i, token, test_ctx.token_to_piece(token, false).c_str());539        if (!test_ctx.decode_token(token, 0)) {540            GGML_ASSERT(false && "Failed to decode token");541        }542    }543 544    printf("min-p sampling test PASSED\n");545}546 547static void test_backend_top_p_sampling(const test_params & params) {548    const int seq_id = 0;549    const float p = 0.9;550    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();551    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));552    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_top_p(p, 0));553    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};554 555    test_context test_ctx(params, backend_sampler_configs);556 557    if (!test_ctx.decode({{seq_id, "Hello"}})) {558        return;559    }560 561    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);562 563    float * logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx);564    uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);565 566    // Print the logits that are above the min-p threshold567    std::vector<float> filtered_logits;568    for (size_t i = 0; i < n_logits; ++i) {569        if (logits[i] > -1e9f) {570            filtered_logits.push_back(logits[i]);571        }572    }573    GGML_ASSERT(filtered_logits.size() < (size_t) test_ctx.n_vocab);574    GGML_ASSERT(filtered_logits.size() > 0);575 576    // Sample using CPU sampler for verification to inspect they are reasonable577    struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();578    llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));579    llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(88));580 581    llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);582    const std::string token_str = test_ctx.token_to_piece(token, false);583    printf("top-p cpu sampled token id:%d, string: '%s'\n", token, token_str.c_str());584    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);585 586    // Decode and sample 10 more tokens587    for (int i = 0; i < 10; i++) {588        int32_t loop_idx = test_ctx.idx_for_seq(seq_id);589        llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), loop_idx);590        printf("top-p gen step %d: token id :%5.d, string: %s\n", i, token, test_ctx.token_to_piece(token, false).c_str());591        test_ctx.decode_token(token, 0);592    }593 594    printf("top-p sampling test PASSED\n");595}596 597static void test_backend_multi_sequence_sampling(const test_params & params) {598    struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();599    llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));600    llama_sampler_chain_add(sampler_chain_0.get(), llama_sampler_init_greedy());601 602    struct llama_sampler_chain_params chain_params_1 = llama_sampler_chain_default_params();603    llama_sampler_ptr sampler_chain_1(llama_sampler_chain_init(chain_params_1));604    llama_sampler_chain_add(sampler_chain_1.get(), llama_sampler_init_temp(0.8f));605    llama_sampler_chain_add(sampler_chain_1.get(), llama_sampler_init_greedy());606 607    std::vector<llama_sampler_seq_config> backend_sampler_configs = {608        { 0, sampler_chain_0.get() },609        { 1, sampler_chain_1.get() }610    };611 612    test_context test_ctx(params, backend_sampler_configs);613 614    std::map<llama_seq_id, std::string> prompts = {615        {0, "Hello"},616        {1, "Some"}617    };618 619    if (!test_ctx.decode(prompts)) {620        GGML_ASSERT(false && "Failed to decode token");621    }622 623    // Verify sequence 0624    {625        int32_t batch_idx = test_ctx.idx_for_seq(0);626        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);627        const std::string token_str = test_ctx.token_to_piece(token, false);628        printf("Seq 0 sampled token id=%d, string='%s'\n", token, token_str.c_str());629        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);630    }631 632    // Verify sequence 1633    {634        int32_t batch_idx= test_ctx.idx_for_seq(1);635        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);636        const std::string token_str = test_ctx.token_to_piece(token, false);637        printf("Seq 1 sampled token id=%d, string='%s'\n", token, token_str.c_str());638        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);639    }640 641    // Generate tokens for each sequence642    printf("\nMulti-sequence generation:\n");643    for (int step = 0; step < 4; step++) {644        std::map<llama_seq_id, llama_token> tokens;645 646        for (llama_seq_id seq_id : {0, 1}) {647            int32_t idx = test_ctx.idx_for_seq(seq_id);648            llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), idx);649            const std::string token_str = test_ctx.token_to_piece(token, false);650            printf("  Seq %d, step %d: token id=%d, string='%s'\n", seq_id, step, token, token_str.c_str());651            tokens[seq_id] = token;652        }653 654        // Decode all tokens in a single batch655        if (!test_ctx.decode_tokens(tokens)) {656            GGML_ASSERT(false && "Failed to decode token");657        }658    }659 660    printf("backend multi-sequence sampling test PASSED\n");661}662 663static void test_backend_dist_sampling(const test_params & params) {664    const int seq_id = 189;665    const int32_t seed = 88;666 667    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();668    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));669    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_dist(seed));670    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};671 672    test_context test_ctx(params, backend_sampler_configs);673 674    if (!test_ctx.decode({{seq_id, "Some"}})) {675        GGML_ASSERT(false && "Failed to decode token");676    }677 678    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);679    llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);680    printf("dist sampled id:%d, string:'%s'\n", token, test_ctx.token_to_piece(token, false).c_str());681    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);682    //GGML_ASSERT(llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx) == nullptr);683 684    token = llama_get_sampled_token_ith(test_ctx.ctx.get(), -1);685    printf("dist sampled id:%d, string:'%s'\n", token, test_ctx.token_to_piece(token, false).c_str());686    GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);687 688    printf("backend dist sampling test PASSED\n");689}690 691static void test_backend_dist_sampling_and_cpu(const test_params & params) {692    const int seq_id = 0;693    const int32_t seed = 88;694 695    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();696    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));697    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_dist(seed));698    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};699 700    test_context test_ctx(params, backend_sampler_configs);701 702    if (!test_ctx.decode({{seq_id, "Some"}})) {703        GGML_ASSERT(false && "Failed to decode token");704    }705 706    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);707 708    // Sample using CPU sampler709    struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();710    llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));711    llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(18));712 713    llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);714    llama_token cpu_token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);715    printf("dist & cpu sampled id:%d, string:'%s'\n", cpu_token, test_ctx.token_to_piece(cpu_token, false).c_str());716    GGML_ASSERT(backend_token == cpu_token);717 718    printf("backend dist & cpu sampling test PASSED\n");719}720 721static void test_backend_logit_bias_sampling(const test_params & params) {722    const auto * model = params.model.get();723    const auto * vocab = llama_model_get_vocab(model);724 725    const int seq_id = 0;726 727    std::vector<llama_logit_bias> logit_bias;728 729    // Get the token for the piece "World".730    const std::string piece = "World";731    std::vector<llama_token> tokens(16);732    llama_tokenize(vocab, piece.c_str(), piece.size(), tokens.data(), tokens.size(), false, false);733 734    llama_token bias_token = tokens[0];735    // TODO: biasing too much here makes the Vulkan sampling fail - should be investigated further736    //       https://github.com/ggml-org/llama.cpp/actions/runs/20894267644/job/60030252675?pr=18753#step:3:23350737    //logit_bias.push_back({ bias_token, +100.0f });738    logit_bias.push_back({ bias_token, +10.0f });739 740    printf("biasing token piece '%s' -> token id %d\n", piece.c_str(), bias_token);741 742    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();743    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));744    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_logit_bias(745                llama_vocab_n_tokens(vocab),746                logit_bias.size(),747                logit_bias.data()));748    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_dist(88));749 750    std::vector<llama_sampler_seq_config> backend_sampler_configs = {751        { seq_id, backend_sampler_chain.get() },752    };753 754    test_context test_ctx(params, backend_sampler_configs);755 756    if (!test_ctx.decode({{seq_id, "Hello"}})) {757        GGML_ASSERT(false && "Failed to decode token");758    }759 760    llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), test_ctx.idx_for_seq(seq_id));761    printf("sampled token = %d, expected = %d\n", backend_token, bias_token);762    GGML_ASSERT(backend_token == bias_token);763 764    printf("backend logit bias sampling test PASSED\n");765}766 767// This test verifies that it is possible to have two different backend samplers,768// one that uses the backend dist sampler, and another that uses CPU dist sampler.769static void test_backend_mixed_sampling(const test_params & params) {770    struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();771    llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));772    llama_sampler_chain_add(sampler_chain_0.get(), llama_sampler_init_dist(88));773 774    int k = 40;775    struct llama_sampler_chain_params chain_params_1 = llama_sampler_chain_default_params();776    llama_sampler_ptr sampler_chain_1(llama_sampler_chain_init(chain_params_1));777    llama_sampler_chain_add(sampler_chain_1.get(), llama_sampler_init_top_k(k));778 779    std::vector<llama_sampler_seq_config> backend_sampler_configs = {780        { 0, sampler_chain_0.get() },781        { 1, sampler_chain_1.get() }782    };783 784    test_context test_ctx(params, backend_sampler_configs);785 786    std::map<llama_seq_id, std::string> prompts = {787        {0, "Hello"},788        {1, "Some"}789    };790 791    if (!test_ctx.decode(prompts)) {792        GGML_ASSERT(false && "Failed to decode token");793    }794 795    // Verify sequence 0 that used the dist backend sampler.796    {797        int32_t batch_idx = test_ctx.idx_for_seq(0);798        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);799        const std::string token_str = test_ctx.token_to_piece(token, false);800        printf("sampled token id=%d, string='%s'\n", token, token_str.c_str());801        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);802        //GGML_ASSERT(llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx) == nullptr);803        //GGML_ASSERT(llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx) == 0);804    }805 806    // Verify sequence 1 that used the top-k backend sampler.807    {808        int32_t batch_idx = test_ctx.idx_for_seq(1);809        float * logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), batch_idx);810        GGML_ASSERT(logits != nullptr);811        size_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), batch_idx);812        GGML_ASSERT(n_logits == (size_t) k);813        GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx) == LLAMA_TOKEN_NULL);814    }815 816    printf("backend mixed sampling test PASSED\n");817}818 819static void test_backend_set_sampler(const test_params & params) {820    const int seq_id = 0;821    const int32_t seed = 88;822 823    struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();824    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));825    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_dist(seed));826    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};827 828    test_context test_ctx(params, backend_sampler_configs);829 830    if (!test_ctx.decode({{seq_id, "Hello"}})) {831        GGML_ASSERT(false && "Failed to decode token");832    }833 834    int32_t batch_idx = test_ctx.idx_for_seq(seq_id);835 836    // Sample using backend sampler configured above837    llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);838    const std::string backend_token_str = test_ctx.token_to_piece(backend_token, false);839    printf("dist sampled token = %d, string='%s'\n", backend_token, backend_token_str.c_str());840 841    // Now clear the backend sampler for this sequence.842    llama_set_sampler(test_ctx.ctx.get(), seq_id, nullptr);843    printf("Cleared backend sampler for seq_id %d\n", seq_id);844 845    // Sample using CPU sampler846    struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();847    llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));848    llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(18));849 850    std::map<llama_seq_id, llama_token> tokens = { { seq_id, backend_token}, };851    if (!test_ctx.decode_tokens(tokens)) {852        GGML_ASSERT(false && "Failed to decode token");853    }854 855    // Should not have any sampled token or probs after clearing the backend sampler.856    const int32_t idx = test_ctx.idx_for_seq(seq_id);857    GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), idx) == LLAMA_TOKEN_NULL);858    GGML_ASSERT(llama_get_sampled_probs_ith(test_ctx.ctx.get(), idx) == nullptr);859 860    // Sample the token using the CPU sampler chain.861    llama_token token2 = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), seq_id);862    const std::string token2_str = test_ctx.token_to_piece(token2, false);863    printf("CPU sampled token after clearing backend sampler: id=%d, string='%s'\n", token2, token2_str.c_str());864    std::map<llama_seq_id, llama_token> tokens2 = { { seq_id, token2}, };865 866    // Set a new backend sampler for the sequence.867    struct llama_sampler_chain_params new_backend_chain_params = llama_sampler_chain_default_params();868    llama_sampler_ptr new_backend_sampler_chain(llama_sampler_chain_init(new_backend_chain_params));869    llama_sampler_chain_add(new_backend_sampler_chain.get(), llama_sampler_init_top_k(20));870    llama_sampler_chain_add(new_backend_sampler_chain.get(), llama_sampler_init_dist(seed));871    llama_set_sampler(test_ctx.ctx.get(), seq_id, new_backend_sampler_chain.get());872 873    if (!test_ctx.decode_tokens(tokens2)) {874        GGML_ASSERT(false && "Failed to decode token");875    }876 877    llama_token new_backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), test_ctx.idx_for_seq(seq_id));878    const std::string new_backend_token_str = test_ctx.token_to_piece(new_backend_token, false);879    printf("dist sampled token = %d, string='%s'\n", new_backend_token, new_backend_token_str.c_str());880 881    printf("backend set sampler test PASSED\n");882}883 884static void test_backend_cpu_mixed_batch(const test_params & params) {885    // Sequence 0 uses backend sampling886    struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();887    llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));888    llama_sampler_chain_add(sampler_chain_0.get(), llama_sampler_init_dist(88));889 890    std::vector<llama_sampler_seq_config> backend_sampler_configs = {891        { 0, sampler_chain_0.get() },892    };893 894    // We need 2 sequences: seq 0 with backend sampling, seq 1 with CPU sampling895    test_context test_ctx(params, backend_sampler_configs, 2);896 897    std::map<llama_seq_id, std::string> prompts = {898        {0, "Hello"}, // Will use backend sampling899        {1, "Some"}   // Will use CPU sampling900    };901 902    if (!test_ctx.decode(prompts)) {903        GGML_ASSERT(false && "Failed to decode token");904    }905 906    // Verify sequence 0 (backend sampled)907    {908        int32_t batch_idx = test_ctx.idx_for_seq(0);909        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);910        const std::string token_str = test_ctx.token_to_piece(token, false);911        printf("Seq 0 (backend) sampled token id=%d, string='%s'\n", token, token_str.c_str());912        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);913    }914 915    // Verify sequence 1 (CPU sampled)916    {917        int32_t batch_idx = test_ctx.idx_for_seq(1);918 919        llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);920        GGML_ASSERT(backend_token == LLAMA_TOKEN_NULL);921 922        struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();923        llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));924        llama_sampler_chain_add(chain.get(), llama_sampler_init_greedy());925 926        llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);927        const std::string token_str = test_ctx.token_to_piece(token, false);928        printf("Seq 1 (CPU) sampled token id=%d, string='%s'\n", token, token_str.c_str());929        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);930    }931 932    // Clear/remove the backend sampler, and sample again933    {934        // clear the backend sampler for seq 0 so that there are no backend935        // samplers.936        llama_set_sampler(test_ctx.ctx.get(), 0, nullptr);937 938        // Create a CPU sampler and verify we can sample from it.939        struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();940        llama_sampler_ptr chain(llama_sampler_chain_init(chain_params));941        llama_sampler_chain_add(chain.get(), llama_sampler_init_greedy());942 943        int32_t batch_idx = test_ctx.idx_for_seq(1);944        llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), batch_idx);945        if (!test_ctx.decode_token(token, 1)) {946            GGML_ASSERT(false && "Failed to decode token");947        }948    }949 950    // Set a backend sampler so that we can verify that it can be reset951    {952        struct llama_sampler_chain_params chain_params = llama_sampler_chain_default_params();953        llama_sampler_ptr sampler_chain(llama_sampler_chain_init(chain_params));954        llama_sampler_chain_add(sampler_chain.get(), llama_sampler_init_dist(88));955 956        llama_set_sampler(test_ctx.ctx.get(), 0, sampler_chain.get());957 958        if (!test_ctx.decode_token(3834, 0)) {959            GGML_ASSERT(false && "Failed to decode token");960        }961 962        int32_t batch_idx = test_ctx.idx_for_seq(0);963        llama_token token = llama_get_sampled_token_ith(test_ctx.ctx.get(), batch_idx);964        const std::string token_str = test_ctx.token_to_piece(token, false);965        printf("re-added backend sampled token id=%d, string='%s'\n", token, token_str.c_str());966        GGML_ASSERT(token >= 0 && token < test_ctx.n_vocab);967    }968 969    printf("backend-cpu mixed batch test PASSED\n");970}971 972static void test_backend_max_outputs(const test_params & params) {973    const int seq_id = 0;974    const int32_t seed = 88;975 976    llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();977    llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(backend_chain_params));978    llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_dist(seed));979    std::vector<llama_sampler_seq_config> backend_sampler_configs = {{ seq_id, backend_sampler_chain.get() }};980 981    test_context test_ctx(params, backend_sampler_configs);982 983    llama_batch batch = llama_batch_init(512, 0, 1);984    std::string prompt = "Hello";985 986    std::vector<llama_token> tokens;987    tokens.push_back(llama_vocab_bos(test_ctx.vocab));988 989    std::vector<llama_token> prompt_tokens(32);990    int n_tokens = llama_tokenize(test_ctx.vocab, prompt.c_str(), prompt.length(),991                                   prompt_tokens.data(), prompt_tokens.size(),992                                   false, false);993    for (int i = 0; i < n_tokens; i++) {994        tokens.push_back(prompt_tokens[i]);995    }996 997    for (size_t i = 0; i < tokens.size(); i++) {998        // set all tokens as output to trigger error999        common_batch_add(batch, tokens[i], i, { seq_id }, true);1000    }1001 1002    printf(">>> test_max_outputs expected error start:\n");1003    const int ret = llama_decode(test_ctx.ctx.get(), batch);1004    GGML_ASSERT(ret != 0 && "llama_decode should not succeed multiple outputs per sequence");1005    printf("<<< test_max_outputs expected error end.\n");1006    llama_batch_free(batch);1007 1008    printf("backend max outputs test PASSED\n");1009}1010 1011struct backend_test_case {1012    std::string name;1013    void (*fn)(const test_params &);1014    bool enabled_by_default;1015};1016 1017static const backend_test_case BACKEND_TESTS[] = {1018    { "greedy",          test_backend_greedy_sampling,         true  },1019    { "logit_bias",      test_backend_logit_bias_sampling,     true  },1020    { "temp",            test_backend_temp_sampling,           true  },1021    { "temp_ext",        test_backend_temp_ext_sampling,       true  },1022    { "top_k",           test_backend_top_k_sampling,          true  },1023    { "multi_sequence",  test_backend_multi_sequence_sampling, true  },1024    { "dist",            test_backend_dist_sampling,           true  },1025    { "dist_and_cpu",    test_backend_dist_sampling_and_cpu,   true  },1026    { "set_sampler",     test_backend_set_sampler,             true  },1027    { "max_outputs",     test_backend_max_outputs,             true  },1028    { "mixed",           test_backend_mixed_sampling,          true  },1029    { "min_p",           test_backend_min_p_sampling,          true  },1030    { "cpu_mixed",       test_backend_cpu_mixed_batch,         true  },1031    { "top_p",           test_backend_top_p_sampling,          true  },1032};1033 1034static test_args parse_cli(int argc, char ** argv) {1035    test_args out;1036 1037    for (int i = 1; i < argc; ++i) {1038        const char * arg = argv[i];1039 1040        if (std::strcmp(arg, "--test") == 0) {1041            if (i + 1 >= argc) {1042                fprintf(stderr, "--test expects a value\n");1043                exit(EXIT_FAILURE);1044            }1045            out.test = argv[++i];1046            continue;1047        }1048        if (std::strncmp(arg, "--test=", 7) == 0) {1049            out.test = arg + 7;1050            continue;1051        }1052        if (std::strcmp(arg, "--model") == 0) {1053            if (i + 1 >= argc) {1054                fprintf(stderr, "--model expects a value\n");1055                exit(EXIT_FAILURE);1056            }1057            out.model = argv[++i];1058            continue;1059        }1060        if (std::strncmp(arg, "--model=", 8) == 0) {1061            out.model = arg + 8;1062            continue;1063        }1064        if (std::strcmp(arg, "--device") == 0) {1065            if (i + 1 >= argc) {1066                fprintf(stderr, "--device expects a value (cpu or gpu)\n");1067                exit(EXIT_FAILURE);1068            }1069            out.device = argv[++i];1070            continue;1071        }1072        if (std::strncmp(arg, "--device=", 9) == 0) {1073            out.device = arg + 9;1074            continue;1075        }1076        if (out.model.empty()) {1077            out.model = arg;1078            continue;1079        }1080 1081        fprintf(stderr, "Unexpected argument: %s\n", arg);1082        exit(EXIT_FAILURE);1083    }1084 1085    if (out.device != "cpu" && out.device != "gpu" && out.device != "auto") {1086        fprintf(stderr, "Invalid device '%s'. Must be 'cpu', 'gpu' or 'auto'\n", out.device.c_str());1087        exit(EXIT_FAILURE);1088    }1089 1090    return out;1091}1092 1093static std::vector<const backend_test_case *> collect_tests_to_run(const std::string & requested) {1094    std::vector<const backend_test_case *> selected;1095 1096    if (!requested.empty()) {1097        for (const auto & test : BACKEND_TESTS) {1098            if (test.name == requested) {1099                selected.push_back(&test);1100                break;1101            }1102        }1103        if (selected.empty()) {1104            fprintf(stderr, "Unknown test '%s'. Available tests:\n", requested.c_str());1105            for (const auto & test : BACKEND_TESTS) {1106                fprintf(stderr, "  %s\n", test.name.c_str());1107            }1108            exit(EXIT_FAILURE);1109        }1110    } else {1111        for (const auto & test : BACKEND_TESTS) {1112            if (test.enabled_by_default) {1113                selected.push_back(&test);1114            }1115        }1116    }1117 1118    if (selected.empty()) {1119        fprintf(stderr, "No backend sampling tests selected. Use --test=<name> to pick one.\n");1120    }1121 1122    return selected;1123}1124 1125static void run_tests(const std::vector<const backend_test_case *> & tests, const test_params & args) {1126    for (const auto & test : tests) {1127        fprintf(stderr, "\n=== %s ===\n", test->name.c_str());1128        try {1129            test->fn(args);1130        } catch (const std::exception & e) {1131            fprintf(stderr, "Error running test '%s': %s\n", test->name.c_str(), e.what());1132            exit(EXIT_FAILURE);1133        }1134    }1135}1136 1137int main(int argc, char ** argv) {1138    test_args args = parse_cli(argc, argv);1139 1140    if (args.model.empty()) {1141        args.model = get_model_or_exit(1, argv);1142    }1143 1144    {1145        std::ifstream file(args.model);1146        if (!file.is_open()) {1147            fprintf(stderr, "no model '%s' found\n", args.model.c_str());1148            return EXIT_FAILURE;1149        }1150    }1151 1152    fprintf(stderr, "using '%s'\n", args.model.c_str());1153 1154    llama_backend_init();1155 1156    test_params params = {1157        /*.model =*/ load_model(args),1158    };1159 1160    const std::vector<const backend_test_case *> tests = collect_tests_to_run(args.test);1161    if (!tests.empty()) {1162        run_tests(tests, params);1163    }1164 1165    return 0;1166}1167