| 377 | } |
| 378 | |
| 379 | llama_token llama_sampler_sample(struct llama_sampler * smpl, struct llama_context * ctx, int32_t idx) { |
| 380 | const auto * logits = llama_get_logits_ith(ctx, idx); |
| 381 | |
| 382 | const llama_model * model = llama_get_model(ctx); |
| 383 | const llama_vocab * vocab = llama_model_get_vocab(model); |
| 384 | |
| 385 | const int n_vocab = llama_vocab_n_tokens(vocab); |
| 386 | |
| 387 | // TODO: do not allocate each time |
| 388 | std::vector<llama_token_data> cur; |
| 389 | cur.reserve(n_vocab); |
| 390 | for (llama_token token_id = 0; token_id < n_vocab; token_id++) { |
| 391 | cur.emplace_back(llama_token_data{token_id, logits[token_id], 0.0f}); |
| 392 | } |
| 393 | |
| 394 | llama_token_data_array cur_p = { |
| 395 | /* .data = */ cur.data(), |
| 396 | /* .size = */ cur.size(), |
| 397 | /* .selected = */ -1, |
| 398 | /* .sorted = */ false, |
| 399 | }; |
| 400 | |
| 401 | llama_sampler_apply(smpl, &cur_p); |
| 402 | |
| 403 | GGML_ASSERT(cur_p.selected >= 0 && cur_p.selected < (int32_t) cur_p.size); |
| 404 | |
| 405 | auto token = cur_p.data[cur_p.selected].id; |
| 406 | |
| 407 | llama_sampler_accept(smpl, token); |
| 408 | |
| 409 | return token; |
| 410 | } |
| 411 | |
| 412 | // sampler chain |
| 413 |
no test coverage detected