| 389 | } |
| 390 | |
| 391 | llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t n_seqs) { |
| 392 | const uint32_t n_tokens = n_seq_tokens*n_seqs; |
| 393 | |
| 394 | clear(); |
| 395 | split_reset(); |
| 396 | |
| 397 | auto udata = std::make_shared<llama_ubatch::data_t>(); |
| 398 | |
| 399 | udata->token .resize(n_tokens); |
| 400 | udata->embd .clear(); |
| 401 | udata->pos .resize(n_tokens); |
| 402 | udata->n_seq_id .resize(n_tokens); |
| 403 | udata->seq_id .resize(n_tokens); |
| 404 | udata->seq_id_unq.resize(0); |
| 405 | udata->seq_idx .resize(LLAMA_MAX_SEQ, -1); |
| 406 | udata->output .resize(n_tokens); |
| 407 | |
| 408 | for (uint32_t s = 0; s < n_seqs; ++s) { |
| 409 | udata->seq_idx[s] = s; |
| 410 | udata->seq_id_unq.push_back(s); |
| 411 | } |
| 412 | |
| 413 | llama_ubatch res { |
| 414 | /*.b_equal_seqs =*/ true, |
| 415 | /*.n_tokens =*/ n_tokens, |
| 416 | /*.n_seq_tokens =*/ n_seq_tokens, |
| 417 | /*.n_seqs =*/ n_seqs, |
| 418 | /*.n_seqs_unq =*/ n_seqs, |
| 419 | /*.n_pos =*/ n_pos_per_embd, |
| 420 | |
| 421 | /*.token =*/ udata->token.data(), |
| 422 | /*.embd =*/ nullptr, |
| 423 | /*.pos =*/ udata->pos.data(), |
| 424 | /*.n_seq_id =*/ udata->n_seq_id.data(), |
| 425 | /*.seq_id =*/ udata->seq_id.data(), |
| 426 | /*.seq_id_unq =*/ udata->seq_id_unq.data(), |
| 427 | /*.seq_idx =*/ udata->seq_idx.data(), |
| 428 | /*.output =*/ udata->output.data(), |
| 429 | /*.data =*/ std::move(udata), |
| 430 | }; |
| 431 | |
| 432 | return res; |
| 433 | } |
| 434 | |
| 435 | const llama_batch & llama_batch_allocr::get_batch() const { |
| 436 | return batch; |
no test coverage detected