mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 02:47:26 -05:00
add server code
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
#include "server-context.h"
|
||||
#include "server-chat.h"
|
||||
#include "server-common.h"
|
||||
#include "server-decision.h"
|
||||
#include "server-http.h"
|
||||
#include "server-task.h"
|
||||
#include "server-queue.h"
|
||||
@@ -825,6 +826,8 @@ public:
|
||||
mtmd_helper_init_opt init_opt = mtmd_helper_init_opt_default();
|
||||
const llama_vocab * vocab = nullptr;
|
||||
|
||||
server_decision_context decision;
|
||||
|
||||
server_queue queue_tasks;
|
||||
server_response queue_results;
|
||||
|
||||
@@ -1102,6 +1105,13 @@ private:
|
||||
|
||||
vocab = llama_model_get_vocab(model_tgt);
|
||||
|
||||
try {
|
||||
decision.init(model_tgt);
|
||||
} catch (const std::exception & e) {
|
||||
SRV_ERR("failed to init decision model: %s\n", e.what());
|
||||
return false;
|
||||
}
|
||||
|
||||
n_ctx = llama_n_ctx(ctx_tgt);
|
||||
|
||||
add_bos_token = llama_vocab_get_add_bos(vocab);
|
||||
@@ -2194,6 +2204,49 @@ private:
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
void send_decision(const server_slot & slot, const common_batch & batch, int32_t i_batch) {
|
||||
auto res = std::make_unique<server_task_result_decision>();
|
||||
res->id = slot.task->id;
|
||||
res->index = slot.task->index;
|
||||
res->n_tokens = slot.task->n_tokens();
|
||||
|
||||
const auto & decision = slot.task->decision;
|
||||
|
||||
if (!decision.labels.empty()) {
|
||||
const float * logits = llama_get_logits_ith(slot.ctx_tgt, i_batch);
|
||||
if (logits == nullptr) {
|
||||
send_error(slot, "failed to get logits", ERROR_TYPE_SERVER);
|
||||
return;
|
||||
}
|
||||
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
|
||||
for (const llama_token label : decision.labels) {
|
||||
GGML_ASSERT(label >= 0 && label < n_vocab);
|
||||
res->scores.push_back(logits[label]);
|
||||
}
|
||||
} else {
|
||||
// the prompt is evaluated in one batch, the n-th output of this slot is the n-th prompt token
|
||||
std::vector<int32_t> idx;
|
||||
for (int i = 0; i < batch.size(); ++i) {
|
||||
if (batch.tokens[i].output && batch.tokens[i].seq_id == slot.id) {
|
||||
idx.push_back(i);
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(decision.column >= 0 && decision.column < llama_model_n_embd_out(model_tgt));
|
||||
for (const int32_t marker : decision.markers) {
|
||||
const float * embd = marker >= 0 && marker < (int32_t) idx.size() ? llama_get_embeddings_ith(slot.ctx_tgt, idx[marker]) : nullptr;
|
||||
if (embd == nullptr) {
|
||||
send_error(slot, "failed to get embeddings", ERROR_TYPE_SERVER);
|
||||
return;
|
||||
}
|
||||
res->scores.push_back(embd[decision.column]);
|
||||
}
|
||||
}
|
||||
|
||||
SLT_DBG(slot, "%s", "sending decision result\n");
|
||||
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
void send_rerank(const server_slot & slot, const common_batch & batch) {
|
||||
auto res = std::make_unique<server_task_result_rerank>();
|
||||
res->id = slot.task->id;
|
||||
@@ -2384,6 +2437,7 @@ private:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
case SERVER_TASK_TYPE_EMBEDDING:
|
||||
case SERVER_TASK_TYPE_RERANK:
|
||||
case SERVER_TASK_TYPE_DECISION:
|
||||
{
|
||||
// special case: if input is provided via CLI, tokenize it first
|
||||
// otherwise, no need to tokenize as it's already done inside the HTTP thread
|
||||
@@ -3831,6 +3885,13 @@ private:
|
||||
return;
|
||||
}
|
||||
|
||||
if (slot.task->type == SERVER_TASK_TYPE_DECISION) {
|
||||
send_decision(slot, batch.view, slot.i_batch - off);
|
||||
slot.release();
|
||||
slot.i_batch = -1;
|
||||
return;
|
||||
}
|
||||
|
||||
GGML_ASSERT(slot.task->need_sampling());
|
||||
|
||||
// prompt evaluated for next-token prediction
|
||||
@@ -5224,6 +5285,64 @@ void server_routes::init_routes() {
|
||||
return res;
|
||||
};
|
||||
|
||||
this->post_systemone = [this](const server_http_req & req) {
|
||||
auto res = create_response();
|
||||
const auto & decision = ctx_server.decision;
|
||||
if (decision.type == SERVER_DECISION_TYPE_NONE) {
|
||||
res->error(format_error_response("This model is not a decision model", ERROR_TYPE_NOT_SUPPORTED));
|
||||
return res;
|
||||
}
|
||||
if (decision.need_embd() && (!params.embedding || meta->pooling_type != LLAMA_POOLING_TYPE_NONE)) {
|
||||
res->error(format_error_response("This decision model requires `--embedding --pooling none`", ERROR_TYPE_NOT_SUPPORTED));
|
||||
return res;
|
||||
}
|
||||
|
||||
const json body = json::parse(req.body);
|
||||
const auto questions = decision.parse_questions(body);
|
||||
|
||||
// one task per question
|
||||
auto & rd = res->rd;
|
||||
{
|
||||
std::vector<server_task> tasks;
|
||||
tasks.reserve(questions.size());
|
||||
for (const auto & question : questions) {
|
||||
server_task task = server_task(SERVER_TASK_TYPE_DECISION);
|
||||
task.id = rd.get_new_id();
|
||||
decision.fill_task(body.at("state"), question, task);
|
||||
tasks.push_back(std::move(task));
|
||||
}
|
||||
rd.post_tasks(std::move(tasks));
|
||||
}
|
||||
|
||||
auto all_results = rd.wait_for_all(req.should_stop);
|
||||
|
||||
if (all_results.is_terminated) {
|
||||
return res; // connection is closed
|
||||
} else if (all_results.error) {
|
||||
res->error(all_results.error->to_json());
|
||||
return res;
|
||||
}
|
||||
|
||||
json answers = json::object();
|
||||
int32_t n_tokens = 0;
|
||||
for (size_t i = 0; i < questions.size(); 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);
|
||||
n_tokens += result->n_tokens;
|
||||
}
|
||||
|
||||
res->ok(json{
|
||||
{"model", meta->model_name},
|
||||
{"answers", answers},
|
||||
{"usage", {
|
||||
{"input_tokens", n_tokens},
|
||||
{"output_tokens", 0},
|
||||
}},
|
||||
});
|
||||
return res;
|
||||
};
|
||||
|
||||
this->get_lora_adapters = [this](const server_http_req & req) {
|
||||
auto res = create_response();
|
||||
|
||||
|
||||
Reference in New Issue
Block a user