Compare commits

...

1 Commits

Author SHA1 Message Date
Xuan Son Nguyen f45ca59f0f server: abstract llama_memory calls to common_memory 2026-07-28 11:36:16 +02:00
3 changed files with 58 additions and 48 deletions
+29 -3
View File
@@ -1518,23 +1518,49 @@ done:
return res;
}
void common_context_seq_rm(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
static void common_context_seq_rm(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
auto * mem = llama_get_memory(ctx);
if (!llama_memory_seq_rm(mem, seq_id, p0, p1)) {
GGML_ABORT("%s", string_format("failed to remove sequence %d with p0=%d, p1=%d\n", seq_id, p0, p1).c_str());
}
}
void common_context_seq_cp(llama_context * ctx, llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
static void common_context_seq_cp(llama_context * ctx, llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
auto * mem = llama_get_memory(ctx);
llama_memory_seq_cp(mem, seq_id_src, seq_id_dst, p0, p1);
}
void common_context_seq_add(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta) {
static void common_context_seq_add(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta) {
auto * mem = llama_get_memory(ctx);
llama_memory_seq_add(mem, seq_id, p0, p1, delta);
}
void common_memory::init(llama_context * ctx_tgt, llama_context * ctx_dft) {
this->ctx_tgt = ctx_tgt;
this->ctx_dft = ctx_dft;
}
void common_memory::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) const {
common_context_seq_rm(ctx_tgt, seq_id, p0, p1);
if (ctx_dft) {
common_context_seq_rm(ctx_dft, seq_id, p0, p1);
}
}
void common_memory::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) const {
common_context_seq_cp(ctx_tgt, seq_id_src, seq_id_dst, p0, p1);
if (ctx_dft) {
common_context_seq_cp(ctx_dft, seq_id_src, seq_id_dst, p0, p1);
}
}
void common_memory::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta) const {
common_context_seq_add(ctx_tgt, seq_id, p0, p1, delta);
if (ctx_dft) {
common_context_seq_add(ctx_dft, seq_id, p0, p1, delta);
}
}
void common_set_adapter_lora(struct llama_context * ctx, std::vector<common_adapter_lora_info> & lora) {
std::vector<llama_adapter_lora *> loras;
std::vector<float> scales;
+11 -4
View File
@@ -948,10 +948,17 @@ enum common_context_seq_rm_type {
// note: clears the memory of the context
common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx);
// aborts execution on failure
void common_context_seq_rm (llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1);
void common_context_seq_add(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta);
void common_context_seq_cp (llama_context * ctx, llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1);
struct common_memory {
llama_context * ctx_tgt = nullptr;
llama_context * ctx_dft = nullptr;
void init(llama_context * ctx_tgt, llama_context * ctx_dft = nullptr);
// aborts execution on failure
void seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) const;
void seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta) const;
void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) const;
};
//
// Batch utils
+18 -41
View File
@@ -164,6 +164,8 @@ struct server_slot {
llama_context * ctx_tgt = nullptr;
llama_context * ctx_dft = nullptr;
common_memory mem;
// multimodal
mtmd_context * mctx = nullptr;
mtmd::batch_ptr mbatch = nullptr;
@@ -253,10 +255,7 @@ struct server_slot {
void prompt_clear() {
SLT_TRC(*this, "clearing prompt with %zu tokens\n", prompt.tokens.size());
common_context_seq_rm(ctx_tgt, id, -1, -1);
if (ctx_dft) {
common_context_seq_rm(ctx_dft, id, -1, -1);
}
mem.seq_rm(id, -1, -1);
prompt.clear();
}
@@ -668,13 +667,8 @@ struct server_slot {
void copy_state_to(server_slot & other) const {
GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT);
common_context_seq_rm(ctx_tgt, other.id, -1, -1);
common_context_seq_cp(ctx_tgt, id, other.id, -1, -1);
if (ctx_dft) {
common_context_seq_rm(ctx_dft, other.id, -1, -1);
common_context_seq_cp(ctx_dft, id, other.id, -1, -1);
}
mem.seq_rm(other.id, -1, -1);
mem.seq_cp(id, other.id, -1, -1);
other.n_decoded = n_decoded;
other.n_remaining = n_remaining;
@@ -1302,6 +1296,7 @@ private:
slot.id = i;
slot.ctx_tgt = ctx_tgt;
slot.ctx_dft = ctx_dft;
slot.mem.init(ctx_tgt, ctx_dft);
slot.spec = spec.get();
slot.n_ctx = n_ctx_slot;
@@ -2881,13 +2876,8 @@ private:
SLT_WRN(slot, "slot context shift, n_keep = %d, n_left = %d, n_discard = %d\n", n_keep, n_left, n_discard);
common_context_seq_rm (ctx_tgt, slot.id, n_keep , n_keep + n_discard);
common_context_seq_add(ctx_tgt, slot.id, n_keep + n_discard, slot.prompt.n_tokens(), -n_discard);
if (ctx_dft) {
common_context_seq_rm (ctx_dft, slot.id, n_keep , n_keep + n_discard);
common_context_seq_add(ctx_dft, slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -n_discard);
}
slot.mem.seq_rm (slot.id, n_keep , n_keep + n_discard);
slot.mem.seq_add(slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -n_discard);
// add generated tokens to cache
// ref: https://github.com/ggml-org/llama.cpp/pull/16818#discussion_r2473269481
@@ -2998,7 +2988,9 @@ private:
ckpt.load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
}
common_context_seq_rm(ctx_dft, slot.id, ckpt.pos_max + 1, -1);
if (!llama_memory_seq_rm(llama_get_memory(ctx_dft), slot.id, ckpt.pos_max + 1, -1)) {
GGML_ABORT("failed to remove sequence %d\n", slot.id);
}
}
if (!draft.empty()) {
@@ -3201,13 +3193,8 @@ private:
const int64_t kv_shift = (int64_t) head_p - (int64_t) head_c;
common_context_seq_rm (ctx_tgt, slot.id, head_p, head_c);
common_context_seq_add(ctx_tgt, slot.id, head_c, head_c + n_match, kv_shift);
if (ctx_dft) {
common_context_seq_rm (ctx_dft, slot.id, head_p, head_c);
common_context_seq_add(ctx_dft, slot.id, head_c, head_c + n_match, kv_shift);
}
slot.mem.seq_rm (slot.id, head_p, head_c);
slot.mem.seq_add(slot.id, head_c, head_c + n_match, kv_shift);
for (size_t i = 0; i < n_match; i++) {
slot.prompt.tokens.set_token(head_p + i, slot.prompt.tokens[head_c + i]);
@@ -3379,10 +3366,7 @@ private:
SLT_TRC(slot, "cached n_tokens = %d, memory_seq_rm [%d, end)\n", slot.prompt.n_tokens(), p0);
common_context_seq_rm(ctx_tgt, slot.id, p0, -1);
if (ctx_dft) {
common_context_seq_rm(ctx_dft, slot.id, p0, -1);
}
slot.mem.seq_rm(slot.id, p0, -1);
// If using an alora, there may be uncached tokens that come
// before the invocation sequence. When this happens, the
@@ -3837,18 +3821,14 @@ private:
SLT_DBG(slot, "restoring speculative checkpoint (pos_min = %d, pos_max = %d, size = %zu)\n", ckpt.pos_min, ckpt.pos_max, ckpt.size());
{
ckpt.load_tgt(slot.ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
common_context_seq_rm(slot.ctx_tgt, slot.id, ckpt.pos_max + 1, -1);
}
ckpt.load_tgt(slot.ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
if (slot.ctx_dft) {
ckpt.load_dft(slot.ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
common_context_seq_rm(slot.ctx_dft, slot.id, ckpt.pos_max + 1, -1);
}
slot.mem.seq_rm(slot.id, ckpt.pos_max + 1, -1);
slot.prompt.tokens.keep_first(ckpt.n_tokens);
slot.smpl = std::move(smpl_save);
@@ -3889,10 +3869,7 @@ private:
slot.sampled = ids.back(); // last accepted token
SLT_DBG(slot, "add accepted tokens: sampled=%d, ids.size=%zu, n_draft=%zu\n", slot.sampled, ids.size(), n_draft);
common_context_seq_rm(slot.ctx_tgt, slot.id, slot.prompt.tokens.pos_next(), -1);
if (slot.ctx_dft) {
common_context_seq_rm(slot.ctx_dft, slot.id, slot.prompt.tokens.pos_next(), -1);
}
slot.mem.seq_rm(slot.id, slot.prompt.tokens.pos_next(), -1);
for (size_t i = 0; i < ids.size(); ++i) {
completion_token_output result;