mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 03:17:32 -05:00
init support for clef (text only)
This commit is contained in:
@@ -168,6 +168,16 @@ bool llama_batch_allocr::init(
|
||||
}
|
||||
}
|
||||
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
if (batch_inp.tokens[i].decision_order != 0) {
|
||||
decision_order.resize(n_tok, 0);
|
||||
break;
|
||||
}
|
||||
}
|
||||
for (size_t i = 0; i < decision_order.size(); ++i) {
|
||||
decision_order[i] = batch_inp.tokens[i].decision_order;
|
||||
}
|
||||
|
||||
//
|
||||
// set up the internal llama_batch to point to our owned arrays
|
||||
//
|
||||
@@ -765,6 +775,7 @@ void llama_batch_allocr::clear() {
|
||||
seq_id .clear();
|
||||
seq_id_unq .clear();
|
||||
output .clear();
|
||||
decision_order.clear();
|
||||
|
||||
for (auto & cur : seq_pos) {
|
||||
cur.clear();
|
||||
@@ -826,6 +837,10 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
|
||||
udata->n_seq_id[i] = batch.n_seq_id[idxs[i]];
|
||||
udata->output[i] = batch.logits[idxs[i]];
|
||||
|
||||
if (!decision_order.empty()) {
|
||||
udata->decision_order.push_back(decision_order[idxs[i]]);
|
||||
}
|
||||
|
||||
for (int s = 0; s < udata->n_seq_id[i]; ++s) {
|
||||
const llama_seq_id seq_id = batch.seq_id[idxs[i]][s];
|
||||
|
||||
@@ -870,6 +885,10 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
|
||||
/*.data =*/ std::move(udata),
|
||||
};
|
||||
|
||||
if (!res.data->decision_order.empty()) {
|
||||
res.decision_order = res.data->decision_order.data();
|
||||
}
|
||||
|
||||
if (debug > 0) {
|
||||
LLAMA_LOG_DEBUG("%s: added ubatch to split:\n", __func__);
|
||||
|
||||
@@ -1175,6 +1194,14 @@ bool llama_batch_ext::set_output(int32_t idx, bool output_last) {
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_batch_ext::set_decision_order(int32_t idx, int32_t order) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
tokens[idx].decision_order = order;
|
||||
return true;
|
||||
}
|
||||
|
||||
// llama_batch_ext C API
|
||||
|
||||
llama_batch_ext * llama_batch_ext_init(llama_context * ctx) {
|
||||
@@ -1243,6 +1270,10 @@ bool llama_batch_ext_set_output_logits(llama_batch_ext * batch, int32_t idx, boo
|
||||
return batch->set_output(idx, value);
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_decision_order(llama_batch_ext * batch, int32_t idx, int32_t order) {
|
||||
return batch->set_decision_order(idx, order);
|
||||
}
|
||||
|
||||
// llama_batch_compat
|
||||
|
||||
void llama_batch_compat::init(llama_batch_ext & dst, const llama_batch & batch_inp, size_t n_embd_row) {
|
||||
|
||||
Reference in New Issue
Block a user