decode a batch of tokens by evaluating the transformer - lctx: llama context - batch: batch to evaluate return 0 on success return positive int on warning return negative int on error
| 6762 | // return negative int on error |
| 6763 | // |
| 6764 | static int llama_decode_internal( |
| 6765 | llama_context & lctx, |
| 6766 | llama_batch batch) { |
| 6767 | const uint32_t n_tokens = batch.n_tokens; |
| 6768 | |
| 6769 | if (n_tokens == 0) { |
| 6770 | LLAMA_LOG_ERROR("%s: n_tokens == 0", __func__); |
| 6771 | return -1; |
| 6772 | } |
| 6773 | |
| 6774 | const auto & model = lctx.model; |
| 6775 | const auto & hparams = model.hparams; |
| 6776 | const auto & cparams = lctx.cparams; |
| 6777 | |
| 6778 | const auto n_batch = cparams.n_batch; |
| 6779 | |
| 6780 | GGML_ASSERT(n_tokens <= n_batch); |
| 6781 | |
| 6782 | int n_threads = n_tokens == 1 ? cparams.n_threads : cparams.n_threads_batch; |
| 6783 | GGML_ASSERT((!batch.token && batch.embd) || (batch.token && !batch.embd)); // NOLINT |
| 6784 | |
| 6785 | const int64_t t_start_us = ggml_time_us(); |
| 6786 | |
| 6787 | #ifdef GGML_USE_MPI |
| 6788 | // TODO: needs fix after #3228 |
| 6789 | GGML_ASSERT(false && "not implemented"); |
| 6790 | //ggml_mpi_eval_init(lctx.ctx_mpi, &n_tokens, &n_past, &n_threads); |
| 6791 | #endif |
| 6792 | |
| 6793 | GGML_ASSERT(n_threads > 0); |
| 6794 | |
| 6795 | auto & kv_self = lctx.kv_self; |
| 6796 | |
| 6797 | GGML_ASSERT(!!kv_self.ctx); |
| 6798 | |
| 6799 | const int64_t n_embd = hparams.n_embd; |
| 6800 | const int64_t n_vocab = hparams.n_vocab; |
| 6801 | |
| 6802 | // helpers for smoother batch API transistion |
| 6803 | // after deprecating the llama_eval calls, these will be removed |
| 6804 | std::vector<llama_pos> pos; |
| 6805 | |
| 6806 | std::vector<int32_t> n_seq_id; |
| 6807 | std::vector<llama_seq_id *> seq_id_arr; |
| 6808 | std::vector<std::vector<llama_seq_id>> seq_id; |
| 6809 | |
| 6810 | if (batch.pos == nullptr) { |
| 6811 | pos.resize(n_tokens); |
| 6812 | for (uint32_t i = 0; i < n_tokens; i++) { |
| 6813 | pos[i] = batch.all_pos_0 + i*batch.all_pos_1; |
| 6814 | } |
| 6815 | |
| 6816 | batch.pos = pos.data(); |
| 6817 | } |
| 6818 | |
| 6819 | if (batch.seq_id == nullptr) { |
| 6820 | n_seq_id.resize(n_tokens); |
| 6821 | seq_id.resize(n_tokens); |
no test coverage detected