Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 14 additions & 1 deletion tools/server/server-context.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2247,6 +2247,14 @@ struct server_context_impl {
void create_checkpoint(server_slot & slot, const int64_t n_tokens_cur, llama_pos pos_min, llama_pos pos_max) {
const int id_task = slot.task->id;

// the buffers of an evicted checkpoint go to the new one: a fresh vector is zero-filled by resize()
std::vector<uint8_t> spare_tgt;
std::vector<uint8_t> spare_dft;
const auto keep_spare = [&](common_prompt_checkpoint & ckpt) {
spare_tgt = std::move(ckpt.data_tgt);
spare_dft = std::move(ckpt.data_dft);
};

// evict checkpoints within min-step of a previous checkpoint, unless they were
// created by the current task
int64_t last = -1;
Expand All @@ -2255,6 +2263,7 @@ struct server_context_impl {
SLT_TRC(slot, "erasing context checkpoint too close to an earlier one (pos_min = %d, pos_max = %d, n_tokens = %" PRId64 ", size = %.3f MiB)\n",
it->pos_min, it->pos_max, it->n_tokens, (float) it->size() / 1024 / 1024);

keep_spare(*it);
it = slot.prompt.checkpoints.erase(it);
continue;
}
Expand All @@ -2265,18 +2274,22 @@ struct server_context_impl {

while (slot.prompt.checkpoints.size() >= (size_t) params_base.n_ctx_checkpoints) {
// make room for the new checkpoint, if needed
const auto & cur = slot.prompt.checkpoints.front();
auto & cur = slot.prompt.checkpoints.front();

SLT_WRN(slot, "erasing old context checkpoint (pos_min = %d, pos_max = %d, n_tokens = %" PRId64 ", size = %.3f MiB)\n",
cur.pos_min, cur.pos_max, cur.n_tokens, (float) cur.size() / 1024 / 1024);

keep_spare(cur);
slot.prompt.checkpoints.erase(slot.prompt.checkpoints.begin());
}

auto & cur = slot.prompt.checkpoints.emplace_back();

cur.id_task = id_task;

cur.data_tgt = std::move(spare_tgt);
cur.data_dft = std::move(spare_dft);

// [TAG_CHECKPOINTS_FIX_POS_MIN]
// TODO: here we incorrectly deterimne that the saved checkpoint data covers the [pos_min, pos_max] range
// this is not true for SWA models: https://github.com/ggml-org/llama.cpp/pull/24411#issuecomment-4677983225
Expand Down