server: abstract llama_memory calls to common_memory (#26221)
This commit is contained in:
+29
-3
@@ -1518,23 +1518,49 @@ done:
|
|||||||
return res;
|
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);
|
auto * mem = llama_get_memory(ctx);
|
||||||
if (!llama_memory_seq_rm(mem, seq_id, p0, p1)) {
|
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());
|
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);
|
auto * mem = llama_get_memory(ctx);
|
||||||
llama_memory_seq_cp(mem, seq_id_src, seq_id_dst, p0, p1);
|
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);
|
auto * mem = llama_get_memory(ctx);
|
||||||
llama_memory_seq_add(mem, seq_id, p0, p1, delta);
|
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) {
|
void common_set_adapter_lora(struct llama_context * ctx, std::vector<common_adapter_lora_info> & lora) {
|
||||||
std::vector<llama_adapter_lora *> loras;
|
std::vector<llama_adapter_lora *> loras;
|
||||||
std::vector<float> scales;
|
std::vector<float> scales;
|
||||||
|
|||||||
+10
-3
@@ -949,10 +949,17 @@ enum common_context_seq_rm_type {
|
|||||||
// note: clears the memory of the context
|
// note: clears the memory of the context
|
||||||
common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx);
|
common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx);
|
||||||
|
|
||||||
|
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
|
// aborts execution on failure
|
||||||
void common_context_seq_rm (llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1);
|
void seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) const;
|
||||||
void common_context_seq_add(llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta);
|
void seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos delta) const;
|
||||||
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);
|
void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) const;
|
||||||
|
};
|
||||||
|
|
||||||
//
|
//
|
||||||
// Batch utils
|
// Batch utils
|
||||||
|
|||||||
@@ -164,6 +164,8 @@ struct server_slot {
|
|||||||
llama_context * ctx_tgt = nullptr;
|
llama_context * ctx_tgt = nullptr;
|
||||||
llama_context * ctx_dft = nullptr;
|
llama_context * ctx_dft = nullptr;
|
||||||
|
|
||||||
|
common_memory mem;
|
||||||
|
|
||||||
// multimodal
|
// multimodal
|
||||||
mtmd_context * mctx = nullptr;
|
mtmd_context * mctx = nullptr;
|
||||||
mtmd::batch_ptr mbatch = nullptr;
|
mtmd::batch_ptr mbatch = nullptr;
|
||||||
@@ -253,10 +255,7 @@ struct server_slot {
|
|||||||
void prompt_clear() {
|
void prompt_clear() {
|
||||||
SLT_TRC(*this, "clearing prompt with %zu tokens\n", prompt.tokens.size());
|
SLT_TRC(*this, "clearing prompt with %zu tokens\n", prompt.tokens.size());
|
||||||
|
|
||||||
common_context_seq_rm(ctx_tgt, id, -1, -1);
|
mem.seq_rm(id, -1, -1);
|
||||||
if (ctx_dft) {
|
|
||||||
common_context_seq_rm(ctx_dft, id, -1, -1);
|
|
||||||
}
|
|
||||||
|
|
||||||
prompt.clear();
|
prompt.clear();
|
||||||
}
|
}
|
||||||
@@ -668,13 +667,8 @@ struct server_slot {
|
|||||||
void copy_state_to(server_slot & other) const {
|
void copy_state_to(server_slot & other) const {
|
||||||
GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT);
|
GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT);
|
||||||
|
|
||||||
common_context_seq_rm(ctx_tgt, other.id, -1, -1);
|
mem.seq_rm(other.id, -1, -1);
|
||||||
common_context_seq_cp(ctx_tgt, id, other.id, -1, -1);
|
mem.seq_cp(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);
|
|
||||||
}
|
|
||||||
|
|
||||||
other.n_decoded = n_decoded;
|
other.n_decoded = n_decoded;
|
||||||
other.n_remaining = n_remaining;
|
other.n_remaining = n_remaining;
|
||||||
@@ -1302,6 +1296,7 @@ private:
|
|||||||
slot.id = i;
|
slot.id = i;
|
||||||
slot.ctx_tgt = ctx_tgt;
|
slot.ctx_tgt = ctx_tgt;
|
||||||
slot.ctx_dft = ctx_dft;
|
slot.ctx_dft = ctx_dft;
|
||||||
|
slot.mem.init(ctx_tgt, ctx_dft);
|
||||||
slot.spec = spec.get();
|
slot.spec = spec.get();
|
||||||
slot.n_ctx = n_ctx_slot;
|
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);
|
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);
|
slot.mem.seq_rm (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);
|
slot.mem.seq_add(slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -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);
|
|
||||||
}
|
|
||||||
|
|
||||||
// add generated tokens to cache
|
// add generated tokens to cache
|
||||||
// ref: https://github.com/ggml-org/llama.cpp/pull/16818#discussion_r2473269481
|
// 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);
|
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()) {
|
if (!draft.empty()) {
|
||||||
@@ -3201,13 +3193,8 @@ private:
|
|||||||
|
|
||||||
const int64_t kv_shift = (int64_t) head_p - (int64_t) head_c;
|
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);
|
slot.mem.seq_rm (slot.id, head_p, head_c);
|
||||||
common_context_seq_add(ctx_tgt, slot.id, head_c, head_c + n_match, kv_shift);
|
slot.mem.seq_add(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);
|
|
||||||
}
|
|
||||||
|
|
||||||
for (size_t i = 0; i < n_match; i++) {
|
for (size_t i = 0; i < n_match; i++) {
|
||||||
slot.prompt.tokens.set_token(head_p + i, slot.prompt.tokens[head_c + 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);
|
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);
|
slot.mem.seq_rm(slot.id, p0, -1);
|
||||||
if (ctx_dft) {
|
|
||||||
common_context_seq_rm(ctx_dft, slot.id, p0, -1);
|
|
||||||
}
|
|
||||||
|
|
||||||
// If using an alora, there may be uncached tokens that come
|
// If using an alora, there may be uncached tokens that come
|
||||||
// before the invocation sequence. When this happens, the
|
// 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());
|
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);
|
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);
|
|
||||||
}
|
|
||||||
|
|
||||||
if (slot.ctx_dft) {
|
if (slot.ctx_dft) {
|
||||||
ckpt.load_dft(slot.ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
|
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.prompt.tokens.keep_first(ckpt.n_tokens);
|
||||||
slot.smpl = std::move(smpl_save);
|
slot.smpl = std::move(smpl_save);
|
||||||
|
|
||||||
@@ -3889,10 +3869,7 @@ private:
|
|||||||
slot.sampled = ids.back(); // last accepted token
|
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);
|
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);
|
slot.mem.seq_rm(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);
|
|
||||||
}
|
|
||||||
|
|
||||||
for (size_t i = 0; i < ids.size(); ++i) {
|
for (size_t i = 0; i < ids.size(); ++i) {
|
||||||
completion_token_output result;
|
completion_token_output result;
|
||||||
|
|||||||
Reference in New Issue
Block a user