echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
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 