Brunobkr/llama.cpp_AlgMor24_github
ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.
03k
1#include "server-context.h"2#include "server-chat.h"3#include "server-common.h"4#include "server-http.h"5#include "server-task.h"6#include "server-queue.h"7#include "server-schema.h"8#include "server-stream.h"9 10#include "build-info.h"11#include "common.h"12#include "fit.h"13#include "llama.h"14#include "log.h"15#include "sampling.h"16#include "speculative.h"17#include "mtmd.h"18#include "mtmd-helper.h"19 20#include <algorithm>21#include <cstddef>22#include <cinttypes>23#include <exception>24#include <memory>25#include <filesystem>26#include <utility>27#include <fstream>28 29// fix problem with std::min and std::max30#if defined(_WIN32)31#define WIN32_LEAN_AND_MEAN32#ifndef NOMINMAX33# define NOMINMAX34#endif35#include <windows.h>36#endif37 38using json = nlohmann::ordered_json;39 40constexpr int HTTP_POLLING_SECONDS = 1;41 42static common_speculative_output_limits server_output_limits(const common_params & params) {43 if (params.embedding ||44 (params.pooling_type != LLAMA_POOLING_TYPE_UNSPECIFIED && params.pooling_type != LLAMA_POOLING_TYPE_NONE)) {45 return { params.n_batch, 1 };46 }47 48 auto result = common_speculative_get_output_limits(49 params.n_batch, params.n_parallel, common_speculative_n_max(¶ms.speculative));50 51 result.total = std::max<int32_t>(1, result.total);52 result.per_seq = std::max<int32_t>(1, result.per_seq);53 return result;54}55 56// state diagram: https://github.com/ggml-org/llama.cpp/pull/928357enum slot_state {58 SLOT_STATE_IDLE,59 SLOT_STATE_WAIT_OTHER, // after assigning a task, but waiting for parent slot to process prompt60 SLOT_STATE_STARTED, // after assigning a task and about to process prompt61 SLOT_STATE_PROCESSING_PROMPT,62 SLOT_STATE_DONE_PROMPT,63 SLOT_STATE_GENERATING,64};65 66struct server_slot; // forward declaration67 68struct server_batch {69 llama_batch batch;70 bool batch_rendered = false;71 72 struct token {73 int32_t id_slot;74 llama_token token;75 llama_pos pos;76 bool output;77 };78 std::vector<token> tokens;79 int32_t n_tokens_alloc = 0;80 int32_t n_embd = 0;81 82 // track if given slot can be batched with slots already in the batch83 server_slot * slot_batched = nullptr;84 85 // in embd mode, we temporarily swap out the tokens arr and restore it on clear()86 bool has_embd = false;87 llama_token * tokens_ptr = nullptr;88 std::vector<float> embd;89 90 float alora_scale = -1.0f;91 size_t alora_disabled_id = 0;92 93 server_batch() {94 batch.pos = nullptr; // sentinel: uninitialized batch95 }96 97 ~server_batch() {98 if (batch.pos != nullptr) {99 clear();100 llama_batch_free(batch);101 }102 }103 104 void init(int32_t n_tokens_alloc, int32_t n_embd) {105 this->n_tokens_alloc = n_tokens_alloc;106 this->n_embd = n_embd;107 batch = llama_batch_init(n_tokens_alloc, 0, 1);108 tokens_ptr = batch.token;109 tokens.reserve(n_tokens_alloc);110 }111 112 bool add(int32_t id_slot, llama_token token, llama_pos pos, bool output) {113 GGML_ASSERT(!has_embd); // cannot mix tokens + embd in same batch114 GGML_ASSERT(batch.pos != nullptr);115 if ((int32_t)tokens.size() >= n_tokens_alloc) {116 return false;117 }118 tokens.push_back({ id_slot, token, pos, output });119 return true;120 }121 122 bool add(int32_t id_slot, const std::vector<float> & embd_in, llama_pos pos, bool output) {123 GGML_ASSERT(batch.pos != nullptr);124 if ((int32_t)tokens.size() >= n_tokens_alloc) {125 return false;126 }127 tokens.push_back({ id_slot, LLAMA_TOKEN_NULL, pos, output });128 has_embd = true;129 embd.insert(embd.end(), embd_in.begin(), embd_in.end());130 return true;131 }132 133 void clear() {134 tokens.clear();135 embd.clear();136 common_batch_clear(batch);137 slot_batched = nullptr;138 alora_scale = -1.0f;139 alora_disabled_id = 0;140 batch_rendered = false;141 has_embd = false;142 if (batch.token == nullptr) {143 batch.token = tokens_ptr;144 batch.embd = nullptr;145 }146 }147 148 int32_t size() const {149 return (int32_t)tokens.size();150 }151 152 void set_output(int32_t idx, bool output) {153 GGML_ASSERT(idx >= 0 && idx < (int32_t)tokens.size());154 tokens[idx].output = output;155 }156 157 void render() {158 GGML_ASSERT(!batch_rendered);159 GGML_ASSERT(batch.pos != nullptr);160 common_batch_clear(batch);161 for (int32_t i = 0; i < size(); i++) {162 const auto & t = tokens[i];163 common_batch_add(batch, t.token, t.pos, { t.id_slot }, t.output);164 }165 if (has_embd) {166 batch.token = nullptr; // will be restored on clear()167 batch.embd = embd.data();168 }169 batch_rendered = true;170 }171 172 llama_batch get_view(int32_t off, int32_t n_tokens) const {173 GGML_ASSERT(batch.pos != nullptr);174 GGML_ASSERT(batch_rendered);175 GGML_ASSERT(off >= 0 && off < size());176 GGML_ASSERT(n_tokens > 0 && off + n_tokens <= size());177 178 auto * token = batch.token ? batch.token + off : nullptr;179 auto * embd = batch.embd ? batch.embd + off * n_embd : nullptr;180 181 llama_batch view = {182 n_tokens,183 token,184 embd,185 batch.pos + off,186 batch.n_seq_id + off,187 batch.seq_id + off,188 batch.logits + off,189 };190 191 return view;192 }193};194 195struct server_slot {196 int id;197 198 llama_context * ctx_tgt = nullptr;199 llama_context * ctx_dft = nullptr;200 201 common_memory mem;202 203 // multimodal204 mtmd_context * mctx = nullptr;205 mtmd::batch_ptr mbatch = nullptr;206 207 // speculative decoding208 common_speculative * spec;209 210 llama_tokens spec_draft;211 llama_tokens spec_prompt;212 std::vector<int32_t> spec_i_batch;213 common_prompt_checkpoint spec_ckpt;214 bool spec_is_replay = false;215 216 // TODO: move members that belong to the task (such as `generated_text`, `has_new_line`) to task_results_state217 // see https://github.com/ggml-org/llama.cpp/pull/18283#issuecomment-3710175837218 std::unique_ptr<const server_task> task;219 std::unique_ptr<const server_task> task_prev; // used for debugging220 221 // used to determine the slot that has been used the longest222 int64_t t_last_used = -1;223 224 // generation props225 int32_t n_ctx = 0; // context size per slot226 int32_t n_keep = 0;227 int32_t n_decoded = 0;228 int32_t n_remaining = -1;229 int32_t i_batch = -1;230 231 int32_t n_prompt_tokens_cache = 0;232 int32_t n_prompt_tokens_processed = 0;233 234 size_t last_nl_pos = 0;235 236 std::string generated_text;237 std::string debug_generated_text;238 llama_tokens generated_tokens;239 240 std::vector<completion_token_output> generated_token_probs;241 242 bool has_next_token = true;243 bool has_new_line = false;244 bool truncated = false;245 246 stop_type stop;247 248 std::string stopping_word;249 250 // state251 slot_state state = SLOT_STATE_IDLE;252 253 server_prompt prompt;254 255 bool prompt_save(server_prompt_cache & prompt_cache) const {256 if (prompt.tokens.size() == 0) {257 return false;258 }259 260 const size_t cur_size_tgt = llama_state_seq_get_size_ext(ctx_tgt, id, LLAMA_STATE_SEQ_FLAGS_NONE);261 const size_t cur_size_dft = ctx_dft ? llama_state_seq_get_size_ext(ctx_dft, id, LLAMA_STATE_SEQ_FLAGS_NONE) : 0;262 263 const size_t cur_size = cur_size_tgt + cur_size_dft;264 265 SRV_TRC(" - saving prompt with length %d, total state size = %.3f MiB (draft: %.3f MiB)\n",266 (int) prompt.tokens.size(), cur_size / (1024.0 * 1024.0), cur_size_dft / (1024.0 * 1024.0));267 268 auto * cur = prompt_cache.alloc(prompt, cur_size_tgt, cur_size_dft);269 if (cur == nullptr) {270 return false;271 }272 273 llama_state_seq_get_data_ext(ctx_tgt, cur->data.main.data(), cur_size_tgt, id, LLAMA_STATE_SEQ_FLAGS_NONE);274 if (ctx_dft) {275 llama_state_seq_get_data_ext(ctx_dft, cur->data.drft.data(), cur_size_dft, id, LLAMA_STATE_SEQ_FLAGS_NONE);276 }277 278 return true;279 }280 281 bool prompt_load(server_prompt_cache & prompt_cache, const server_tokens & tokens) {282 bool res = prompt_cache.load(prompt, tokens, ctx_tgt, ctx_dft, id);283 if (!res) {284 SLT_WRN(*this, "%s", "failed to load prompt from cache\n");285 }286 287 return res;288 }289 290 void prompt_clear() {291 SLT_TRC(*this, "clearing prompt with %zu tokens\n", prompt.tokens.size());292 293 mem.seq_rm(id, -1, -1);294 295 prompt.clear();296 }297 298 std::vector<common_adapter_lora_info> lora;299 int32_t alora_invocation_start = -1;300 301 // sampling302 json json_schema;303 304 common_sampler_ptr smpl;305 306 llama_token sampled; // in speculative mode, this is the last accepted token307 308 // for TTS models, this is the embd generated from prev step, decode this to generate next hidden state309 // corresponding to one token position (size = n_embd)310 std::vector<float> inp_embd;311 312 // stats313 size_t n_sent_text = 0; // number of sent text character314 315 // TODO @ngxson : move all metrics to a sub-struct for clarity316 int64_t t_start_process_prompt;317 int64_t t_start_generation;318 int64_t t_print_last = 0;319 int32_t n_decoded_last = 0;320 321 double t_prompt_processing = 0.0; // ms322 double t_token_generation = 0.0; // ms323 324 std::function<void(int /* id_slot */)> callback_on_release;325 326 // Speculative decoding stats327 int32_t n_draft_total = 0; // Total draft tokens generated328 int32_t n_draft_accepted = 0; // Draft tokens actually accepted329 int32_t n_draft_verif_steps = 0; // Total draft token verification steps by the target model330 std::vector<int32_t> n_accepted_per_pos; // Accepted tokens per draft position331 332 void reset() {333 SLT_DBG(*this, "%s", "\n");334 335 spec_is_replay = false;336 337 n_prompt_tokens_cache = 0;338 339 last_nl_pos = 0;340 generated_text = "";341 has_new_line = false;342 truncated = false;343 stop = STOP_TYPE_NONE;344 stopping_word = "";345 n_sent_text = 0;346 347 if (can_speculate()) {348 spec_draft.clear();349 spec_i_batch.clear();350 spec_ckpt.clear();351 }352 generated_tokens.clear();353 generated_token_probs.clear();354 json_schema = json();355 356 // clear speculative decoding stats357 n_draft_total = 0;358 n_draft_accepted = 0;359 n_draft_verif_steps = 0;360 n_accepted_per_pos.clear();361 362 task_prev = std::move(task);363 task.reset();364 365 llama_set_sampler(ctx_tgt, id, nullptr);366 367 // clear alora start368 alora_invocation_start = -1;369 370 // clear multimodal state371 mbatch.reset();372 }373 374 void init_sampler() const {375 common_sampler_reset(smpl.get());376 377 if (!task->need_sampling()) {378 return;379 }380 381 const int64_t t_start = ggml_time_us();382 383 int n_text = 0;384 385 for (int i = 0; i < (int) prompt.tokens.size(); i++) {386 const llama_token id = prompt.tokens[i];387 388 if (id != LLAMA_TOKEN_NULL) {389 common_sampler_accept(smpl.get(), id, false);390 n_text++;391 }392 }393 394 SLT_TRC(*this, "init sampler, took %0.2f ms, tokens: text = %d, total = %d\n",395 (ggml_time_us() - t_start) / 1000.0, n_text, (int) prompt.tokens.size());396 }397 398 bool need_embd() const {399 GGML_ASSERT(task);400 return task->need_embd() || (spec && common_speculative_need_embd(spec));401 }402 403 bool need_embd_nextn() const {404 GGML_ASSERT(task);405 return spec && common_speculative_need_embd_nextn(spec);406 }407 408 // if the context does not have a memory module then all embeddings have to be computed within a single ubatch409 // also we cannot split if the pooling would require any past tokens410 // (MTP supports splitting — uses task->need_embd() not need_embd())411 bool can_split() const {412 GGML_ASSERT(task);413 414 return415 !task->need_embd() ||416 (llama_get_memory(ctx_tgt) && llama_pooling_type(ctx_tgt) == LLAMA_POOLING_TYPE_LAST);417 }418 419 bool can_batch_with(server_slot & other_slot) const {420 GGML_ASSERT(task);421 422 return task->type == other_slot.task->type423 && inp_embd.size() == other_slot.inp_embd.size()424 && are_lora_equal(lora, other_slot.lora);425 }426 427 bool has_budget(const common_params & global_params) {428 GGML_ASSERT(task);429 430 if (task->params.n_predict == -1 && global_params.n_predict == -1) {431 return true; // limitless432 }433 434 n_remaining = -1;435 436 if (task->params.n_predict != -1) {437 n_remaining = task->params.n_predict - n_decoded;438 } else if (global_params.n_predict != -1) {439 n_remaining = global_params.n_predict - n_decoded;440 }441 442 return n_remaining > 0; // no budget443 }444 445 bool is_processing() const {446 return state != SLOT_STATE_IDLE;447 }448 449 bool can_speculate() const {450 return !!spec;451 }452 453 void add_token(const completion_token_output & token) {454 if (!is_processing()) {455 SLT_WRN(*this, "%s", "slot is not processing\n");456 return;457 }458 459 generated_token_probs.push_back(token);460 }461 462 int get_n_draft_max() const {463 GGML_ASSERT(task);464 465 if (!can_speculate()) {466 return 0;467 }468 469 // determine the max draft that fits the current slot state470 // note: slot.prompt is not yet expanded with the `id` token sampled above471 // also, need to leave space for 1 extra token to allow context shifts472 int n_draft_max = n_ctx - prompt.n_tokens() - 2;473 474 if (n_remaining > 0) {475 n_draft_max = std::min(n_draft_max, n_remaining - 1);476 }477 478 SLT_DBG(*this, "max possible draft: %d\n", n_draft_max);479 480 return n_draft_max;481 }482 483 // add sampled token of this slot to the batch, optionally add the speculative draft tokens if any484 void handle_last_sampled_token(server_batch & batch) {485 bool add_ok = true;486 if (spec_draft.empty()) {487 // no speculative decoding488 i_batch = batch.size();489 490 if (!inp_embd.empty()) {491 add_ok &= batch.add(id, inp_embd, prompt.tokens.pos_next(), true);492 } else {493 add_ok &= batch.add(id, sampled, prompt.tokens.pos_next(), true);494 }495 496 SLT_DBG(*this, "slot decode token, id=%d, n_ctx = %d, n_tokens = %d, truncated = %d\n",497 sampled, n_ctx, prompt.n_tokens(), truncated);498 } else {499 SLT_DBG(*this, "generate_draft: id=%d, #tokens=%zu, #draft=%zu, pos_next=%d\n",500 sampled, prompt.tokens.size(), spec_draft.size(), prompt.tokens.pos_next());501 502 GGML_ASSERT(spec_i_batch.empty());503 504 spec_i_batch.push_back(batch.size());505 for (size_t i = 0; i < spec_draft.size(); i++) {506 spec_i_batch.push_back(batch.size() + i + 1);507 }508 509 auto pos0 = prompt.tokens.pos_next();510 511 add_ok &= batch.add(id, sampled, pos0++, true);512 for (auto token : spec_draft) {513 add_ok &= batch.add(this->id, token, pos0++, true);514 }515 }516 517 GGML_ASSERT(add_ok && "batch must be large enough to hold the sampled and draft tokens");518 519 prompt.tokens.push_back(sampled);520 prompt.tokens.insert(spec_draft);521 }522 523 void release() {524 if (is_processing()) {525 GGML_ASSERT(task);526 527 SLT_INF(*this, "stop processing: n_tokens = %d, truncated = %d\n", prompt.n_tokens(), truncated);528 529 t_last_used = ggml_time_us();530 t_token_generation = (ggml_time_us() - t_start_generation) / 1e3;531 532 state = SLOT_STATE_IDLE;533 534 // do not keep context of the child slots - the parent's context is enough535 if (task->is_child()) {536 prompt_clear();537 }538 539 reset();540 541 callback_on_release(id);542 }543 }544 545 result_timings get_timings() const {546 result_timings timings;547 timings.cache_n = n_prompt_tokens_cache;548 549 timings.prompt_n = n_prompt_tokens_processed;550 timings.prompt_ms = t_prompt_processing;551 timings.prompt_per_token_ms = t_prompt_processing / n_prompt_tokens_processed;552 timings.prompt_per_second = 1e3 / t_prompt_processing * n_prompt_tokens_processed;553 554 timings.predicted_n = n_decoded;555 timings.predicted_ms = t_token_generation;556 timings.predicted_per_token_ms = t_token_generation / n_decoded;557 timings.predicted_per_second = 1e3 / t_token_generation * n_decoded;558 559 // Add speculative metrics560 if (n_draft_total > 0) {561 timings.draft_n = n_draft_total;562 timings.draft_n_accepted = n_draft_accepted;563 }564 565 return timings;566 }567 568 size_t find_stopping_strings(const std::string & text, const size_t last_token_size, bool is_full_stop) {569 GGML_ASSERT(task);570 571 size_t stop_pos = std::string::npos;572 573 for (const std::string & word : task->params.antiprompt) {574 size_t pos;575 576 if (is_full_stop) {577 const size_t tmp = word.size() + last_token_size;578 const size_t from_pos = text.size() > tmp ? text.size() - tmp : 0;579 580 pos = text.find(word, from_pos);581 } else {582 // otherwise, partial stop583 pos = string_find_partial_stop(text, word);584 }585 586 if (pos != std::string::npos && (stop_pos == std::string::npos || pos < stop_pos)) {587 if (is_full_stop) {588 stop = STOP_TYPE_WORD;589 stopping_word = word;590 has_next_token = false;591 }592 stop_pos = pos;593 }594 }595 596 return stop_pos;597 }598 599 void print_timings_tg() {600 if (n_decoded < 100) {601 return;602 }603 604 const int64_t t_now = ggml_time_us();605 606 if (t_now - t_print_last < 3*1000*1000) {607 return;608 }609 610 const double n_gen_second = 1e3 / (t_token_generation) * (n_decoded);611 const double n_gen_second_win = 1e6 / (t_now - t_print_last) * (n_decoded - n_decoded_last);612 613 t_print_last = t_now;614 n_decoded_last = n_decoded;615 616 SLT_INF(*this, "n_decoded = %6d, tg = %6.2f t/s, tg_3s = %6.2f t/s\n", n_decoded, n_gen_second, n_gen_second_win);617 }618 619 void print_timings_pp() const {620 const double n_prompt_second = 1e3 / t_prompt_processing * n_prompt_tokens_processed;621 const double f_progress = (float) prompt.n_tokens() / task->n_tokens();622 623 if (t_prompt_processing < 3000.0) {624 return;625 }626 627 SLT_INF(*this, "prompt processing, n_tokens = %6d, progress = %.2f, t = %6.2f s / %.2f tokens per second\n",628 n_prompt_tokens_processed, f_progress, t_prompt_processing / 1e3, n_prompt_second);629 }630 631 void print_timings() const {632 const double t_prompt = t_prompt_processing / n_prompt_tokens_processed;633 const double n_prompt_second = 1e3 / t_prompt_processing * n_prompt_tokens_processed;634 635 const double t_gen = t_token_generation / n_decoded;636 const double n_gen_second = 1e3 / t_token_generation * n_decoded;637 638 SLT_INF(*this,639 "prompt eval time = %10.2f ms / %5d tokens (%8.2f ms per token, %8.2f tokens per second)\n",640 t_prompt_processing, n_prompt_tokens_processed, t_prompt, n_prompt_second);641 642 SLT_INF(*this,643 " eval time = %10.2f ms / %5d tokens (%8.2f ms per token, %8.2f tokens per second)\n",644 t_token_generation, n_decoded, t_gen, n_gen_second);645 646 SLT_INF(*this,647 " total time = %10.2f ms / %5d tokens\n",648 t_prompt_processing + t_token_generation, n_prompt_tokens_processed + n_decoded);649 650 SLT_INF(*this,651 " graphs reused = %10d\n",652 llama_perf_context(ctx_tgt).n_reused);653 654 if (n_draft_total > 0) {655 const float draft_ratio = (float) n_draft_accepted / n_draft_total;656 const double mean_acc_len = n_draft_verif_steps > 0 ? 1.0 + (double) n_draft_accepted / (double) n_draft_verif_steps : 1.0;657 658 std::string acceptance_rates_per_pos;659 if (n_draft_verif_steps > 0) {660 for (size_t i = 0; i < n_accepted_per_pos.size(); ++i) {661 if (i > 0) {662 acceptance_rates_per_pos += ", ";663 }664 acceptance_rates_per_pos += string_format("%.3f", (double) n_accepted_per_pos[i] / (double) n_draft_verif_steps);665 }666 }667 668 SLT_INF(*this,669 "draft acceptance = %0.5f (%5d accepted / %5d generated), mean len = %5.2f\n",670 draft_ratio, n_draft_accepted, n_draft_total, mean_acc_len);671 SLT_TRC(*this,672 " acc per pos = (%s)\n", acceptance_rates_per_pos.c_str());673 }674 675 common_speculative_print_stats(spec);676 }677 678 json to_json(bool only_metrics = false) const {679 json res;680 681 res = {682 {"id", id},683 {"n_ctx", n_ctx},684 {"speculative", can_speculate()},685 {"is_processing", is_processing()},686 };687 688 const auto & ptask = task ? task : task_prev;689 690 if (ptask) {691 res["id_task"] = ptask->id;692 res["n_prompt_tokens"] = (int32_t) prompt.tokens.size();693 res["n_prompt_tokens_processed"] = n_prompt_tokens_processed;694 res["n_prompt_tokens_cache"] = n_prompt_tokens_cache;695 res["params"] = ptask->params.to_json(only_metrics);696 res["next_token"] = {697 {698 {"has_next_token", has_next_token},699 {"has_new_line", has_new_line},700 {"n_remain", n_remaining},701 {"n_decoded", n_decoded},702 }703 };704 705 if (!only_metrics) {706 res["prompt"] = ptask->tokens.detokenize(ctx_tgt, true);707 res["generated"] = generated_text.empty() ? debug_generated_text : generated_text;708 }709 }710 711 return res;712 }713 714 void copy_state_to(server_slot & other) const {715 GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT);716 717 mem.seq_rm(other.id, -1, -1);718 mem.seq_cp(id, other.id, -1, -1);719 720 other.n_decoded = n_decoded;721 other.n_remaining = n_remaining;722 other.i_batch = i_batch;723 724 other.t_start_process_prompt = t_start_process_prompt;725 other.t_prompt_processing = t_prompt_processing;726 other.n_prompt_tokens_cache = n_prompt_tokens_cache;727 other.n_prompt_tokens_processed = n_prompt_tokens_processed;728 729 other.prompt = prompt.clone();730 other.init_sampler();731 }732 733 // returns 0 on success734 // caller need to update prompt.tokens after a successful call to keep track of the processing progress735 int process_mtmd_chunk(size_t idx, size_t & n_tokens_out) {736 GGML_ASSERT(mctx);737 const auto & input_tokens = task->tokens;738 const auto & chunk = input_tokens.find_chunk(idx);739 int32_t res = 0;740 741 auto try_decode = [&]() -> int32_t {742 if (mbatch) {743 float * embd = mtmd_batch_get_output_embd(mbatch.get(), chunk.get());744 if (embd) {745 void * cb_data = spec;746 static auto cb = [](llama_batch batch, void * user_data) {747 common_speculative * spec = static_cast<common_speculative *>(user_data);748 if (!common_speculative_process(spec, batch)) {749 return 1;750 }751 return 0;752 };753 754 llama_pos new_n_past; // unused for now755 res = mtmd_helper_decode_image_chunk(756 mctx,757 ctx_tgt,758 chunk.get(),759 embd,760 prompt.tokens.pos_next(),761 id,762 llama_n_batch(ctx_tgt),763 &new_n_past,764 cb,765 cb_data766 );767 if (res != 0) {768 SLT_ERR(*this, "failed to decode mtmd chunk, idx = %zu, res = %d\n", idx, res);769 return -1;770 }771 n_tokens_out = mtmd_input_chunk_get_n_tokens(chunk.get());772 return 0; // success773 }774 }775 return 1; // (non-error) need to create & encode batch776 };777 778 // if the batch is already exist, try searching & encode779 res = try_decode();780 if (res == 0) {781 return 0;782 }783 if (res < 0) {784 // fatal error785 return res;786 }787 788 // otherwise, the batch is either uninitialized or is used up789 // we need to create & encode a new batch790 mbatch.reset(mtmd_batch_init(mctx));791 res = mtmd_batch_add_chunk(mbatch.get(), chunk.get());792 GGML_ASSERT(res == 0); // we should never have an empty batch793 794 // try batching as much as possible795 int n_added = 1;796 size_t idx_cur = idx;797 while (res == 0) {798 auto [next_chunk, next_idx] = input_tokens.find_next_media_chunk(idx_cur);799 if (next_chunk == nullptr) {800 break;801 }802 res = mtmd_batch_add_chunk(mbatch.get(), next_chunk->get());803 n_added += (res == 0 ? 1 : 0);804 idx_cur = next_idx;805 SLT_DBG(*this, "try adding media chunk idx = %zu to batch, res = %d\n", next_idx, res);806 // if res != 0, batch is full or chunk is not compatible -> this loop breaks807 }808 809 // TODO @ngxson : move this log line to debug when it become more stable810 SLT_TRC(*this, "encoding mtmd batch from idx = %zu, n_chunks = %d\n", idx, n_added);811 812 res = mtmd_batch_encode(mbatch.get());813 if (res != 0) {814 SLT_ERR(*this, "failed to encode mtmd batch for chunk idx = %zu, res = %d\n", idx, res);815 return -1;816 }817 818 return try_decode();819 }820};821 822 823 824//825// server_metrics826//827 828struct server_metrics {829 int64_t t_start = 0;830 831 uint64_t n_prompt_tokens_processed_total = 0;832 uint64_t t_prompt_processing_total = 0;833 uint64_t n_tokens_predicted_total = 0;834 uint64_t t_tokens_generation_total = 0;835 836 uint64_t n_tokens_max = 0;837 838 uint64_t n_prompt_tokens_processed = 0;839 uint64_t t_prompt_processing = 0;840 841 uint64_t n_tokens_predicted = 0;842 uint64_t t_tokens_generation = 0;843 844 uint64_t n_decode_total = 0;845 uint64_t n_busy_slots_total = 0;846 847 uint64_t n_draft_tokens_total = 0;848 uint64_t n_draft_accepted_total = 0;849 uint64_t n_draft_verif_steps_total = 0;850 std::vector<uint64_t> n_accepted_per_pos_total;851 852 void init() {853 t_start = ggml_time_us();854 }855 856 void on_prompt_eval(const server_slot & slot) {857 n_prompt_tokens_processed_total += slot.n_prompt_tokens_processed;858 n_prompt_tokens_processed += slot.n_prompt_tokens_processed;859 t_prompt_processing += slot.t_prompt_processing;860 t_prompt_processing_total += slot.t_prompt_processing;861 862 n_tokens_max = std::max(n_tokens_max, (uint64_t) slot.prompt.n_tokens());863 }864 865 void on_prediction(const server_slot & slot) {866 n_tokens_predicted_total += slot.n_decoded;867 n_tokens_predicted += slot.n_decoded;868 t_tokens_generation += slot.t_token_generation;869 t_tokens_generation_total += slot.t_token_generation;870 871 n_draft_tokens_total += slot.n_draft_total;872 n_draft_accepted_total += slot.n_draft_accepted;873 n_draft_verif_steps_total += slot.n_draft_verif_steps;874 875 if (n_accepted_per_pos_total.size() < slot.n_accepted_per_pos.size()) {876 n_accepted_per_pos_total.resize(slot.n_accepted_per_pos.size(), 0);877 }878 for (size_t i = 0; i < slot.n_accepted_per_pos.size(); i++) {879 n_accepted_per_pos_total[i] += slot.n_accepted_per_pos[i];880 }881 }882 883 void on_decoded(const std::vector<server_slot> & slots) {884 n_decode_total++;885 for (const auto & slot : slots) {886 if (slot.is_processing()) {887 n_busy_slots_total++;888 }889 n_tokens_max = std::max(n_tokens_max, (uint64_t) slot.prompt.n_tokens());890 }891 }892 893 void reset_bucket() {894 n_prompt_tokens_processed = 0;895 t_prompt_processing = 0;896 n_tokens_predicted = 0;897 t_tokens_generation = 0;898 }899};900 901 902//903// server_context_impl (private implementation)904//905 906struct server_context_impl {907 friend struct server_context;908 909public:910 // only use these pointers outside of this class:911 // - when not in sleeping state912 // - and, with thread-safe APIs (e.g., tokenizer calls)913 llama_model * model_tgt = nullptr;914 915 mtmd_context * mctx = nullptr;916 const llama_vocab * vocab = nullptr;917 918 server_queue queue_tasks;919 server_response queue_results;920 921 // note: chat_params must not be refreshed upon existing sleeping state922 server_chat_params chat_params;923 924 server_state_callback_t callback_state = [](server_state, json) -> void {};925 926 server_context_impl() {927 mtmd_helper_log_set(common_log_default_callback, nullptr);928 }929 930 ~server_context_impl() {931 if (!sleeping) {932 // destroy() is already called when entering sleeping state933 // we don't call it again here to avoid double free934 destroy();935 }936 }937 938private:939 // note: accessing these fields outside of this class is not thread-safe940 // use server_context methods instead941 942 common_params params_base;943 944 // note: keep these alive - they determine the lifetime of the model, context, etc.945 common_init_result_ptr llama_init;946 947 llama_context * ctx_tgt = nullptr;948 949 server_batch batch;950 951 llama_model * model_dft = nullptr;952 llama_context * ctx_dft = nullptr;953 954 common_speculative_init_result_ptr spec_init;955 956 common_context_seq_rm_type ctx_tgt_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO;957 common_context_seq_rm_type ctx_dft_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO;958 959 common_speculative_ptr spec;960 961 bool add_bos_token = true;962 963 int32_t n_ctx; // total context for all clients / slots964 965 // set to llama_model_n_swa(model)966 // if swa_full is enabled, this is set to 0 to simulate a non-SWA model967 int32_t n_swa;968 969 // slots / clients970 std::vector<server_slot> slots;971 972 int trace = 0;973 int slots_debug = 0;974 int n_empty_consecutive = 0;975 976 std::unique_ptr<server_prompt_cache> prompt_cache;977 978 server_metrics metrics;979 980 json json_ui_settings = json::object();981 982 // Necessary similarity of prompt for slot selection983 float slot_prompt_similarity = 0.0f;984 985 std::string model_name; // name of the loaded model, to be used by API986 std::set<std::string> model_aliases; // additional names for the model987 std::set<std::string> model_tags; // informational tags988 989 bool sleeping = false;990 991 int64_t t_last_load_progress_ms = 0;992 993 void destroy() {994 spec.reset();995 spec_init.reset();996 997 ctx_dft = nullptr;998 model_dft = nullptr;999 1000 llama_init.reset();1001 1002 ctx_tgt = nullptr;1003 model_tgt = nullptr;1004 1005 mtmd_free(mctx);1006 mctx = nullptr;1007 }1008 1009 void handle_sleeping_state(bool new_state) {1010 GGML_ASSERT(sleeping != new_state);1011 if (new_state) {1012 SRV_INF("%s", "server is entering sleeping state\n");1013 destroy();1014 } else {1015 SRV_INF("%s", "server is exiting sleeping state\n");1016 if (!load_model(params_base)) {1017 GGML_ABORT("failed to reload model after sleeping");1018 }1019 }1020 sleeping = new_state;1021 }1022 1023 struct load_progress_data {1024 server_context_impl * ctx;1025 std::string stage;1026 std::vector<std::string> stages;1027 int64_t t_last_load_progress_ms = 0;1028 load_progress_data(server_context_impl * ctx, const std::string & stage) : ctx(ctx), stage(stage) {}1029 };1030 static bool load_progress_callback(float progress, void * user_data) {1031 auto * d = static_cast<load_progress_data *>(user_data);1032 GGML_ASSERT(d);1033 // always emit the first and final sample; throttle the rest to one per 200ms1034 {1035 auto & t_last = d->t_last_load_progress_ms;1036 const int64_t t_now = ggml_time_ms();1037 const bool first = t_last == 0;1038 const bool done = progress >= 1.0f;1039 const bool throttled = !first && !done && (t_now - t_last) < 200;1040 if (throttled) {1041 return true;1042 }1043 t_last = t_now;1044 }1045 if (d->ctx->callback_state) {1046 d->ctx->callback_state(SERVER_STATE_LOADING, {1047 {"stages", d->stages},1048 {"current", d->stage},1049 {"value", progress},1050 });1051 }1052 return true;1053 }1054 1055 // load the model and initialize llama_context1056 // this may also be called to resume from sleeping state1057 bool load_model(common_params & params) {1058 load_progress_data load_progress_text (this, "text_model");1059 load_progress_data load_progress_mmproj(this, "mmproj_model");1060 load_progress_data load_progress_spec (this, "spec_model");1061 1062 const bool is_resume = sleeping;1063 1064 params_base = params;1065 const auto output_limits = server_output_limits(params_base);1066 params_base.n_outputs_max = output_limits.total;1067 params_base.n_outputs_max_per_seq = output_limits.per_seq;1068 1069 const bool has_mmproj = !params.mmproj.path.empty();1070 const bool has_draft = params.speculative.has_dft();1071 const bool spec_mtp = std::find(params_base.speculative.types.begin(),1072 params_base.speculative.types.end(),1073 COMMON_SPECULATIVE_TYPE_DRAFT_MTP) != params_base.speculative.types.end();1074 const bool has_spec = has_draft || spec_mtp;1075 1076 if (callback_state) {1077 std::vector<std::string> stages = {"text_model"};1078 if (has_spec) {1079 stages.push_back("spec_model");1080 }1081 if (has_mmproj) {1082 stages.push_back("mmproj_model");1083 }1084 load_progress_text.stages = stages;1085 load_progress_mmproj.stages = stages;1086 load_progress_spec.stages = stages;1087 1088 // trigger 0% progress1089 load_progress_callback(0.0f, &load_progress_text);1090 }1091 1092 1093 SRV_INF("loading model '%s'\n", params.model.get_name().c_str());1094 SRV_TRC("local path '%s'\n", params.model.path.c_str());1095 1096 std::string & mmproj_path = params_base.mmproj.path;1097 mtmd_context_params mparams = mtmd_context_params_default();1098 if (has_mmproj) {1099 mparams.use_gpu = params_base.mmproj_use_gpu;1100 mparams.print_timings = false;1101 mparams.n_threads = params_base.cpuparams.n_threads;1102 mparams.flash_attn_type = params_base.flash_attn_type;1103 mparams.warmup = params_base.warmup;1104 mparams.image_min_tokens = params_base.image_min_tokens;1105 mparams.image_max_tokens = params_base.image_max_tokens;1106 mparams.batch_max_tokens = params_base.mtmd_batch_max_tokens;1107 mparams.media_marker = get_media_marker();1108 // progress callback1109 mparams.progress_callback = load_progress_callback;1110 mparams.progress_callback_user_data = &load_progress_mmproj;1111 }1112 1113 // optionally get the memory usage of mmproj1114 if (has_mmproj && params_base.fit_params) {1115 int64_t t_start = ggml_time_us();1116 auto mmproj_mem = mtmd_get_memory_usage(mmproj_path.c_str(), mparams);1117 int64_t t_elapsed = ggml_time_us() - t_start;1118 if (!mmproj_mem.empty()) {1119 size_t total = 0;1120 for (auto & [dev, size] : mmproj_mem) {1121 total += size;1122 }1123 SRV_TRC("[mtmd] estimated worst-case memory usage of mmproj is %.2f MiB (took %.2f ms)\n", total / (1024.0 * 1024.0), t_elapsed / 1000.0);1124 GGML_ASSERT(!params_base.fit_params_target.empty());1125 for (auto & [dev, size] : mmproj_mem) {1126 for (size_t i = 0; i < ggml_backend_dev_count(); i++) {1127 if (ggml_backend_dev_get(i) == dev) {1128 if (i < params_base.fit_params_target.size()) {1129 SRV_DBG("[mtmd] adding %.2f MiB to fit_params_target for device %s\n", size / (1024.0 * 1024.0), ggml_backend_dev_name(dev));1130 params_base.fit_params_target[i] += size;1131 }1132 break;1133 }1134 }1135 }1136 } else {1137 SRV_ERR("%s", "[mtmd] failed to get memory usage of mmproj\n");1138 }1139 }1140 1141 // optionally reserve VRAM for the draft / MTP context before fitting the target model1142 if (params_base.fit_params) {1143 if (has_spec) {1144 // MTP draft context lives on the target model, only context+compute are new1145 bool measure_model_bytes = has_draft;1146 1147 common_params params_dft = common_base_params_to_speculative(params_base);1148 1149 auto mparams_dft = common_model_params_to_llama(params_dft);1150 auto cparams_dft = common_context_params_to_llama(params_dft);1151 if (spec_mtp) {1152 cparams_dft.ctx_type = LLAMA_CONTEXT_TYPE_MTP;1153 }1154 cparams_dft.n_rs_seq = 0;1155 1156 std::vector<ggml_backend_dev_t> devs;1157 uint32_t hp_ngl = 0;1158 uint32_t hp_nct = 0;1159 uint32_t hp_nex = 0;1160 try {1161 auto dmd = common_get_device_memory_data(1162 params_dft.model.path.c_str(), &mparams_dft, &cparams_dft,1163 devs, hp_ngl, hp_nct, hp_nex, GGML_LOG_LEVEL_ERROR);1164 1165 GGML_ASSERT(!params_base.fit_params_target.empty());1166 size_t total = 0;1167 1168 std::vector<ggml_backend_dev_t> tgt_devices = params.devices;1169 1170 if (tgt_devices.empty()) {1171 for(size_t i = 0; i < ggml_backend_dev_count(); ++i) {1172 tgt_devices.push_back(ggml_backend_dev_get(i));1173 }1174 }1175 1176 for (size_t j = 0; j < devs.size(); ++j) {1177 const size_t bytes = (measure_model_bytes ? dmd[j].model : 0) + dmd[j].context + dmd[j].compute;1178 total += bytes;1179 for (size_t i = 0; i < tgt_devices.size(); i++) {1180 if (tgt_devices[i] == devs[j]) {1181 SRV_DBG("[spec] adding %.2f MiB to fit_params_target for device %s\n",1182 bytes / (1024.0 * 1024.0), ggml_backend_dev_name(devs[j]));1183 params_base.fit_params_target[i] += bytes;1184 break;1185 }1186 }1187 }1188 SRV_TRC("[spec] estimated memory usage of %s is %.2f MiB\n",1189 has_draft ? "draft model" : "MTP context",1190 total / (1024.0 * 1024.0));1191 } catch (const std::exception & e) {1192 SRV_WRN("[spec] failed to measure %s memory: %s\n",1193 has_draft ? "draft model" : "MTP context", e.what());1194 }1195 }1196 }1197 1198 // attach a progress callback1199 {1200 params_base.load_progress_callback = load_progress_callback;