mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 19:07:25 -05:00
init support for clef (text only)
This commit is contained in:
@@ -118,6 +118,7 @@ struct server_batch {
|
||||
llama_pos pos;
|
||||
bool output;
|
||||
bool is_prompt; // for stats tracking
|
||||
int32_t decision_order = 0;
|
||||
};
|
||||
std::vector<token> tokens;
|
||||
int32_t n_tokens_alloc = 0;
|
||||
@@ -177,6 +178,11 @@ struct server_batch {
|
||||
tokens[idx].output = output;
|
||||
}
|
||||
|
||||
void set_decision_order(int32_t idx, int32_t order) {
|
||||
GGML_ASSERT(idx >= 0 && idx < (int32_t)tokens.size());
|
||||
tokens[idx].decision_order = order;
|
||||
}
|
||||
|
||||
// render the sub-batch [off, off + n_tokens) into view, index i in view is index off + i here
|
||||
void render(int32_t off, int32_t n_tokens) {
|
||||
GGML_ASSERT(off >= 0 && off < size());
|
||||
@@ -192,6 +198,7 @@ struct server_batch {
|
||||
} else {
|
||||
view.add(t.token, t.pos, t.id_slot, t.output);
|
||||
}
|
||||
view.tokens.back().decision_order = t.decision_order;
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -420,6 +427,11 @@ struct server_slot {
|
||||
bool can_batch_with(server_slot & other_slot) const {
|
||||
GGML_ASSERT(task);
|
||||
|
||||
// a joint decision head reads the whole batch
|
||||
if (!task->decision.order.empty() || !other_slot.task->decision.order.empty()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return task->type == other_slot.task->type
|
||||
&& inp_embd.size() == other_slot.inp_embd.size()
|
||||
&& are_lora_equal(lora, other_slot.lora);
|
||||
@@ -3634,6 +3646,9 @@ private:
|
||||
/* pos = */ slot.prompt.tokens.pos_next(),
|
||||
/* output = */ slot.need_embd(),
|
||||
/* is_prompt = */ true);
|
||||
if (!slot.task->decision.order.empty()) {
|
||||
batch.set_decision_order(batch.size() - 1, slot.task->decision.order[slot.prompt.n_tokens()]);
|
||||
}
|
||||
slot.prompt.tokens.push_back(cur_tok);
|
||||
|
||||
// break at the last user message, or at user messages at least min step past the last checkpoint
|
||||
@@ -5339,12 +5354,21 @@ void server_routes::init_routes() {
|
||||
return res;
|
||||
}
|
||||
|
||||
// one task per question
|
||||
// one task per question, or one task for all of them
|
||||
auto & rd = res->rd;
|
||||
{
|
||||
std::vector<server_task> tasks;
|
||||
tasks.reserve(questions.size());
|
||||
if (decision.is_joint()) {
|
||||
server_task task = server_task(SERVER_TASK_TYPE_DECISION);
|
||||
task.id = rd.get_new_id();
|
||||
decision.fill_task_joint(state, questions, task);
|
||||
tasks.push_back(std::move(task));
|
||||
}
|
||||
for (const auto & question : questions) {
|
||||
if (decision.is_joint()) {
|
||||
break;
|
||||
}
|
||||
server_task task = server_task(SERVER_TASK_TYPE_DECISION);
|
||||
task.id = rd.get_new_id();
|
||||
decision.fill_task(state, question, files, ctx_server.mctx, ctx_server.init_opt, task);
|
||||
@@ -5367,7 +5391,16 @@ void server_routes::init_routes() {
|
||||
|
||||
json answers = json::object();
|
||||
int32_t n_tokens = 0;
|
||||
for (size_t i = 0; i < questions.size(); i++) {
|
||||
if (decision.is_joint()) {
|
||||
auto * result = dynamic_cast<server_task_result_decision *>(all_results.results[0].get());
|
||||
GGML_ASSERT(result != nullptr);
|
||||
const auto scores = decision.split_scores(questions, result->scores);
|
||||
for (size_t i = 0; i < questions.size(); i++) {
|
||||
answers[questions[i].id] = decision.format_answer(questions[i], scores[i]);
|
||||
}
|
||||
n_tokens = result->n_tokens;
|
||||
}
|
||||
for (size_t i = 0; i < questions.size() && !decision.is_joint(); i++) {
|
||||
auto * result = dynamic_cast<server_task_result_decision *>(all_results.results[i].get());
|
||||
GGML_ASSERT(result != nullptr);
|
||||
answers[questions[i].id] = decision.format_answer(questions[i], result->scores);
|
||||
|
||||
Reference in New Issue
Block a user