From bf534d8ba47c4dc101e16d87c210fdea528a89c8 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 5 Sep 2026 11:02:54 +0000 Subject: [PATCH] server: add priority-aware request preemption and slot pause controls Keep live request and speculative state while swapping sequence memory to bounded host RAM. Add client priorities, explicit park and unpark actions, monitoring, and preemption regression tests. --- common/arg.cpp | 10 + common/common.h | 2 + tools/server/README.md | 94 +++++ tools/server/server-common.h | 7 + tools/server/server-context.cpp | 389 +++++++++++++++++- tools/server/server-queue.cpp | 5 +- tools/server/server-schema.cpp | 8 + tools/server/server-task.cpp | 8 + tools/server/server-task.h | 3 + .../tests/unit/test_preempt_priority.py | 198 +++++++++ tools/server/tests/utils.py | 2 +- 11 files changed, 716 insertions(+), 10 deletions(-) create mode 100644 tools/server/tests/unit/test_preempt_priority.py diff --git a/common/arg.cpp b/common/arg.cpp index 86f8610a56d..e1333e6182e 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -3571,6 +3571,16 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.endpoint_slots = value; } ).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_ENDPOINT_SLOTS")); + add_opt(common_arg( + {"--preempt-ram"}, "MiB", + "maximum host RAM for parked sequence snapshots, 0 disables (default: 4096; unified KV only)", + [](common_params & params, int value) { + if (value < 0) { + throw std::invalid_argument("--preempt-ram must be nonnegative"); + } + params.preempt_ram_mib = value; + } + ).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_PREEMPT_RAM")); add_opt(common_arg( {"--slot-save-path"}, "PATH", "path to save slot kv cache (default: disabled)", diff --git a/common/common.h b/common/common.h index de49dac9f63..2e3f8e90038 100644 --- a/common/common.h +++ b/common/common.h @@ -672,6 +672,8 @@ struct common_params { bool log_json = false; + int32_t preempt_ram_mib = 4096; // host sequence snapshots; 0 disables preemption + std::string slot_save_path; std::string media_path; // path to directory for loading media files diff --git a/tools/server/README.md b/tools/server/README.md index 93736c3edfa..fc9dc23af98 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -218,6 +218,7 @@ For the full list of features, please refer to [server's changelog](https://gith | `--metrics` | enable prometheus compatible metrics endpoint (default: disabled)
(env: LLAMA_ARG_ENDPOINT_METRICS) | | `--props` | enable changing global properties via POST /props (default: disabled)
(env: LLAMA_ARG_ENDPOINT_PROPS) | | `--slots, --no-slots` | expose slots monitoring endpoint (default: enabled)
(env: LLAMA_ARG_ENDPOINT_SLOTS) | +| `--preempt-ram MiB` | Maximum RAM for parked target and draft sequence snapshots (default: 4096; 0 disables). Requires unified KV. (env: LLAMA_ARG_PREEMPT_RAM) | | `--slot-save-path PATH` | path to save slot kv cache (default: disabled) | | `--media-path PATH` | directory for loading local media files; files can be accessed via file:// URLs using relative paths (default: disabled) | | `--models-dir PATH` | directory containing models for the router server (default: disabled)
(env: LLAMA_ARG_MODELS_DIR) | @@ -578,6 +579,8 @@ These words will not be included in the completion, so make sure to add them to `id_slot`: Assign the completion task to an specific slot. If is -1 the task will be assigned to a Idle slot. Default: `-1` +`priority`: Signed 32-bit integer scheduling priority, default `0`. Higher values are more important. Supported by native completions and the OpenAI-compatible completion/chat endpoints. See [request preemption](#request-preemption). + `cache_prompt`: Re-use KV cache from a previous request if possible. This way the common prefix does not have to be re-processed, only the suffix that differs between the requests. Because (depending on the backend) the logits are **not** guaranteed to be bit-for-bit identical for different batch sizes (prompt processing vs. token generation) enabling this option can cause nondeterministic results. Default: `true` `return_tokens`: Return the raw generated token ids in the `tokens` field. Otherwise `tokens` remains empty. Default: `false` @@ -1139,6 +1142,97 @@ In *router mode* the query param `?model={model_id}` has to be set. This endpoin | `llamacpp:spec_decode_num_drafts_total` | Counter | Total speculative decoding verification steps (0 when spec-decode is off). | | `llamacpp:spec_decode_num_accepted_tokens_per_pos_total` | Counter | Accepted tokens per draft position (labeled `position="N"`; absent when spec-decode is off or before the first completed speculative request). | +### Request preemption + +With `--kv-unified --preempt-ram 4096`, text requests can share the entire context +pool and pause when its live contents approach capacity. Preemption copies a +sequence's target and draft state to host RAM, removes its device cells, and +keeps its task, sampler, generated tokens, speculative state, reasoning and +stream alive. It restores that state when room returns. No prompt replay or +rollback to a user message is performed. The existing SSE pings continue while +parked; use `sse_ping_interval` to choose the heartbeat interval. + +Scheduling occurs between completed decode/accept steps, before drafting. Each +running sequence reserves one sampled-token cell plus its maximum next draft +(up to three cells with `--spec-type draft-mtp --spec-draft-n-max 2`). The reserve +shrinks with the remaining output/context budget. There is no fixed context +partition or percentage buffer. Prefill is limited to the remaining cells. +Idle cached sequences are cleared before parking live requests. + +On pressure, candidates are the lowest effective-priority tier of independent +slots. The longest resident sequence is protected unless it is the only +candidate in that tier; newer task IDs are selected first among the other +candidates. This gives priority precedence when a long low-priority request is +the last candidate before a more important request. With equal priorities the +longest request survives (ties favor higher effective priority, then slot order). +Parent and child slots from parallel sampling (`n > 1`) are never parked. A pressure survivor +runs until completion before automatic restoration; this avoids repeatedly +swapping at the watermark. Every third park of a task promotes its effective +priority by one, without changing its requested priority. Eligible restores use +highest effective priority, then most parks, then earliest park time. Requests +that do not fit are skipped so a smaller request can resume. + +### POST `/slots/{id_slot}?action=park` and `action=unpark` + +These actions require `--slots`, unified KV, text input and `--preempt-ram > 0`; +they do not require `--slot-save-path`. They execute on the inference task queue, +never during an in-flight decode. For example: + +```sh +curl -X POST 'http://localhost:8080/slots/0?action=park' +curl -X POST 'http://localhost:8080/slots/0?action=unpark' +``` + +`park` holds an active independent slot until `unpark` is requested, including +when other slots are idle. It is idempotent for an already parked slot. `unpark` +clears the manual hold and requests restoration at the next scheduling +opportunity; its response may still say `is_preempted: true`. Both actions return +the slot object, including `priority`, `effective_priority`, `is_preempted`, +`preempt_manual`, `n_preempt`, and `preempt_ram_bytes`. `GET /slots` exposes the +same fields. An idle, invalid, parent or child slot is rejected. A failed manual +snapshot leaves the inference request untouched. + +`--preempt-ram` bounds the additional target/draft snapshot payload, separate +from prompt-cache RAM, samplers, existing checkpoints, generated text, and other +process memory. If the RAM cap or protected slots prevent reclaiming enough +cells, requests remain alive but wait for a cancellation or another action that +releases capacity. Size the cap for the workload: progress cannot be guaranteed +with insufficient host RAM. All-manually-parked or capacity-blocked workloads +wait on the task queue, without a decode busy loop. Cancelling a parked request +frees its snapshot. Park time is excluded from `t_max_predict_ms`. + +Preemption is disabled for non-unified KV and multimodal contexts. Automatic +server sleep (`--sleep-idle-seconds > 0`) is rejected when preemption is enabled, +since unloading a context would invalidate the retained speculative state. +Independent text generation with standard attention and Qwen3.5 MTP is tested; +other speculative architectures and shared-prompt pool exhaustion are not +covered. The original per-request context limit still applies. + +Set `LLAMA_SERVER_PREEMPT_EVERY=N` to exercise a snapshot/remove/restore round +trip whenever a request crosses another N generated tokens. Zero disables the +test hook. With speculation, the crossing can occur up to the accepted draft +length after the exact multiple; the final completed request is not parked. +These forced round trips restore immediately at the same decode boundary. +They isolate snapshot preservation from changes in the concurrent batch. + +State preservation does not guarantee bitwise output identity under arbitrary +CUDA scheduling. Restoring a sequence can change its physical KV placement, +and parking peers changes batch sizes. Floating-point attention reductions and +batch-dependent kernels can then produce different logits and a different +greedy token, even at temperature zero. Byte/token comparisons must distinguish +an isolated forced round trip from a changed concurrent schedule. + +The metrics endpoint additionally exposes: + +| Metric | Type | Description | +| --- | --- | --- | +| `llamacpp:preemptions_total` | Counter | Successful parks, including manual and forced parks. | +| `llamacpp:preempt_restores_total` | Counter | Successful restores. | +| `llamacpp:preempt_blocked_total` | Counter | Scheduling attempts blocked by snapshot RAM or policy. | +| `llamacpp:preempt_copy_seconds_total` | Counter | Time copying snapshots to and from host RAM. | +| `llamacpp:preempt_ram_bytes` | Gauge | Snapshot payload currently held in host RAM. | +| `llamacpp:requests_preempted` | Gauge | Parked requests whose tasks remain alive. | + ### POST `/slots/{id_slot}?action=save`: Save the prompt cache of the specified slot to a file. *Options:* diff --git a/tools/server/server-common.h b/tools/server/server-common.h index f8ea82ef4cf..2da03a6def1 100644 --- a/tools/server/server-common.h +++ b/tools/server/server-common.h @@ -464,6 +464,13 @@ struct server_metrics { uint64_t n_tokens_max = 0; + uint64_t n_preempt = 0; + uint64_t n_restore = 0; + uint64_t n_preempt_blocked = 0; + uint64_t preempt_copy_us = 0; + uint64_t preempt_ram_bytes = 0; + uint64_t n_preempted = 0; + uint64_t n_decode = 0; uint64_t n_busy_slots = 0; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index a9edbd7be8b..ceddb06aea9 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -59,6 +59,7 @@ enum slot_state { SLOT_STATE_PROCESSING_PROMPT, SLOT_STATE_DONE_PROMPT, SLOT_STATE_GENERATING, + SLOT_STATE_PREEMPTED, // task, sampler and stream stay alive; sequence memory is on host }; struct server_slot; // forward declaration @@ -248,6 +249,30 @@ struct server_slot { // state slot_state state = SLOT_STATE_IDLE; + // These objects survive parking, including speculative replay and sampler state. + slot_state preempt_state = SLOT_STATE_IDLE; + std::vector preempt_tgt; + std::vector preempt_dft; + bool preempt_manual = false; + uint64_t n_preempt = 0; + uint64_t preempt_forced_at = 0; + int64_t preempt_since = 0; + int64_t preempt_paused_us = 0; + + size_t preempt_bytes() const { + return preempt_tgt.size() + preempt_dft.size(); + } + + int64_t effective_priority() const { + return task ? (int64_t) task->params.priority + (int64_t) (n_preempt / 3) : 0; + } + + void preempt_clear() { + std::vector().swap(preempt_tgt); + std::vector().swap(preempt_dft); + preempt_manual = false; + } + server_prompt prompt; bool prompt_save(server_prompt_cache & prompt_cache) const { @@ -503,6 +528,11 @@ struct server_slot { t_last_used = ggml_time_us(); + if (state == SLOT_STATE_PREEMPTED) { + // Cancellation must not leave a cache entry for nonresident tokens. + prompt_clear(); + } + preempt_clear(); state = SLOT_STATE_IDLE; // do not keep context of the child slots - the parent's context is enough @@ -645,6 +675,12 @@ struct server_slot { {"n_ctx", n_ctx}, {"speculative", can_speculate()}, {"is_processing", is_processing()}, + {"priority", task ? task->params.priority : (task_prev ? task_prev->params.priority : 0)}, + {"effective_priority", effective_priority()}, + {"is_preempted", state == SLOT_STATE_PREEMPTED}, + {"preempt_manual", preempt_manual}, + {"n_preempt", n_preempt}, + {"preempt_ram_bytes", preempt_bytes()}, }; const auto & ptask = task ? task : task_prev; @@ -863,6 +899,13 @@ struct server_context_impl { int slots_n_diff = 0; // env: LLAMA_SERVER_SLOTS_N_DIFF int n_empty_consecutive = 0; + uint64_t preempt_every = 0; + int preempt_drain_task = -1; // finish this task before automatic restoration + + bool preempt_enabled() const { + return params_base.kv_unified && params_base.preempt_ram_mib > 0 && !mctx && + llama_get_memory(ctx_tgt) != nullptr; + } std::unique_ptr prompt_cache; @@ -1258,6 +1301,17 @@ struct server_context_impl { } } + if (const char * value = getenv("LLAMA_SERVER_PREEMPT_EVERY")) { + const std::string text(value); + if (text.empty() || text.find_first_not_of("0123456789") != std::string::npos) { + throw std::invalid_argument("LLAMA_SERVER_PREEMPT_EVERY must be a nonnegative integer"); + } + preempt_every = std::stoull(text); + } + if (preempt_enabled() && params_base.sleep_idle_seconds > 0) { + throw std::invalid_argument("preemption requires --sleep-idle-seconds disabled (or --preempt-ram 0)"); + } + // the update_slots() logic will always submit a maximum of n_batch or n_parallel tokens // note that n_batch can be > n_ctx (e.g. for non-causal attention models such as BERT where the KV cache is not used) { @@ -1717,6 +1771,9 @@ struct server_context_impl { // the per-request limit takes priority over the global one slot.n_predict_max = task.params.n_predict != -1 ? task.params.n_predict : params_base.n_predict; + slot.n_preempt = 0; + slot.preempt_forced_at = 0; + slot.preempt_paused_us = 0; slot.task = std::make_unique(std::move(task)); slot.state = slot.task->is_child() @@ -1841,7 +1898,7 @@ struct server_context_impl { slot.has_new_line = true; // if we have seen a new line, we stop after a certain time limit, but only upon another new line - if (slot.task->params.t_max_predict_ms > 0 && slot.stats.t_gen_ms() > slot.task->params.t_max_predict_ms) { + if (slot.task->params.t_max_predict_ms > 0 && slot.stats.t_gen_ms() - slot.preempt_paused_us / 1000.0 > slot.task->params.t_max_predict_ms) { slot.stop = STOP_TYPE_LIMIT; slot.has_next_token = false; @@ -2397,6 +2454,11 @@ struct server_context_impl { res->id = task.id; res->n_processing_slots = n_processing_slots; res->n_tasks_deferred = queue_tasks.queue_tasks_deferred_size(); + metrics.preempt_ram_bytes = preempt_ram_used(); + metrics.n_preempted = 0; + for (const auto & slot : slots) { + metrics.n_preempted += slot.state == SLOT_STATE_PREEMPTED; + } res->metrics = metrics; if (task.metrics_reset_bucket) { @@ -2424,6 +2486,34 @@ struct server_context_impl { res->slots_data = std::move(slots_data); res->n_idle_slots = n_idle_slots; + queue_results.send(std::move(res)); + } break; + case SERVER_TASK_TYPE_SLOT_PARK: + case SERVER_TASK_TYPE_SLOT_UNPARK: + { + server_slot * slot = get_slot_by_id(task.slot_action.id_slot); + if (!preempt_enabled()) { + send_error(task, "Parking requires unified KV, text input and --preempt-ram > 0", ERROR_TYPE_NOT_SUPPORTED); + break; + } + if (!slot || !slot->is_processing() || slot->task->is_parent() || slot->task->is_child()) { + send_error(task, "Parking requires an active, independent slot", ERROR_TYPE_INVALID_REQUEST); + break; + } + if (task.type == SERVER_TASK_TYPE_SLOT_PARK) { + if (slot->state != SLOT_STATE_PREEMPTED && !preempt_park(*slot, true)) { + send_error(task, "Unable to park slot within --preempt-ram", ERROR_TYPE_SERVER); + break; + } + slot->preempt_manual = true; + } else { + // Acknowledge the request now; scheduling performs the restore. + slot->preempt_manual = false; + } + auto res = std::make_unique(); + res->id = task.id; + res->slots_data = slot->to_json(true); + res->n_idle_slots = 0; queue_results.send(std::move(res)); } break; case SERVER_TASK_TYPE_SLOT_SAVE: @@ -2674,6 +2764,257 @@ struct server_context_impl { }; #endif + size_t preempt_ram_used() const { + size_t bytes = 0; + for (const auto & slot : slots) { + bytes += slot.preempt_bytes(); + } + return bytes; + } + + // Token positions are a conservative upper bound for full-attention cells. + // In particular, hybrid seq_pos_min() only describes the recurrent rollback + // window and must NOT be used to count the attention KV cells. + int64_t preempt_resident() const { + int64_t cells = 0; + for (const auto & slot : slots) { + if (slot.state != SLOT_STATE_PREEMPTED) { + cells += slot.prompt.tokens.pos_next(); + } + } + return cells; + } + + int64_t preempt_step(const server_slot & slot) const { + const auto state = slot.state == SLOT_STATE_PREEMPTED ? slot.preempt_state : slot.state; + if (state == SLOT_STATE_GENERATING) { + const int n_draft = slot.can_speculate() ? std::max(0, std::min( + common_speculative_n_max(¶ms_base.speculative), slot.get_n_draft_max())) : 0; + return 1 + std::max(n_draft, slot.spec_draft.size()); + } + return state == SLOT_STATE_STARTED || state == SLOT_STATE_PROCESSING_PROMPT ? 1 : 0; + } + + int64_t preempt_demand() const { + int64_t cells = preempt_resident(); + for (const auto & slot : slots) { + if (slot.is_processing() && slot.state != SLOT_STATE_PREEMPTED) { + cells += preempt_step(slot); + } + } + return cells; + } + + bool preempt_candidate(const server_slot & slot) const { + return slot.task && !slot.task->is_parent() && !slot.task->is_child() && + (slot.state == SLOT_STATE_GENERATING || slot.state == SLOT_STATE_PROCESSING_PROMPT || + slot.state == SLOT_STATE_STARTED); + } + + bool preempt_park(server_slot & slot, bool manual) { + if (!preempt_candidate(slot)) { + return false; + } + const int64_t start = ggml_time_us(); + const size_t tgt_size = llama_state_seq_get_size_ext(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_NONE); + const size_t dft_size = ctx_dft ? llama_state_seq_get_size_ext(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_NONE) : 0; + const size_t limit = (size_t) params_base.preempt_ram_mib * 1024 * 1024; + const size_t used = preempt_ram_used(); + if (used > limit || tgt_size > limit - used || dft_size > limit - used - tgt_size) { + return false; + } + try { + std::vector tgt(tgt_size), dft(dft_size); + if (llama_state_seq_get_data_ext(ctx_tgt, tgt.data(), tgt.size(), slot.id, LLAMA_STATE_SEQ_FLAGS_NONE) != tgt.size() || + (ctx_dft && llama_state_seq_get_data_ext(ctx_dft, dft.data(), dft.size(), slot.id, LLAMA_STATE_SEQ_FLAGS_NONE) != dft.size())) { + return false; + } + // No sampler reset, task release, token replay, or stream termination. + // Per-sequence MTP hidden rows and draft samplers remain in `spec`. + slot.preempt_tgt = std::move(tgt); + slot.preempt_dft = std::move(dft); + } catch (const std::bad_alloc &) { + return false; + } + slot.mem.seq_rm(slot.id, -1, -1); + slot.preempt_state = slot.state; + slot.state = SLOT_STATE_PREEMPTED; + slot.preempt_manual = manual; + slot.preempt_since = ggml_time_us(); + ++slot.n_preempt; + ++metrics.n_preempt; + metrics.preempt_copy_us += ggml_time_us() - start; + SLT_INF(slot, "parked: priority=%d, n_preempt=%" PRIu64 ", tokens=%d, ram=%zu, manual=%d\n", + slot.task->params.priority, slot.n_preempt, slot.prompt.n_tokens(), slot.preempt_bytes(), manual); + return true; + } + + bool preempt_restore(server_slot & slot) { + const int64_t start = ggml_time_us(); + const bool tgt_ok = llama_state_seq_set_data_ext(ctx_tgt, slot.preempt_tgt.data(), slot.preempt_tgt.size(), + slot.id, LLAMA_STATE_SEQ_FLAGS_NONE) == slot.preempt_tgt.size(); + const bool dft_ok = tgt_ok && (!ctx_dft || llama_state_seq_set_data_ext(ctx_dft, slot.preempt_dft.data(), + slot.preempt_dft.size(), slot.id, LLAMA_STATE_SEQ_FLAGS_NONE) == slot.preempt_dft.size()); + if (!dft_ok) { + // set_data may have partially populated either context. Keep the host + // snapshot intact and remove the partial restore before retrying. + slot.mem.seq_rm(slot.id, -1, -1); + return false; + } + if (slot.preempt_state == SLOT_STATE_GENERATING) { + slot.preempt_paused_us += ggml_time_us() - slot.preempt_since; + } + slot.state = slot.preempt_state; + slot.preempt_clear(); + ++metrics.n_restore; + metrics.preempt_copy_us += ggml_time_us() - start; + SLT_INF(slot, "restored: n_preempt=%" PRIu64 ", tokens=%d\n", slot.n_preempt, slot.prompt.n_tokens()); + return true; + } + + bool preempt_schedule() { + // Common path: one O(P) scan, no snapshots, sorting or allocations. + if (preempt_every == 0 && preempt_drain_task == -1) { + int64_t demand = 0; + bool parked = false; + bool runnable = false; + for (const auto & slot : slots) { + parked |= slot.state == SLOT_STATE_PREEMPTED; + if (slot.state != SLOT_STATE_PREEMPTED) { + demand += slot.prompt.tokens.pos_next(); + if (slot.is_processing()) { + runnable = true; + demand += preempt_step(slot); + } + } + } + if (!parked && demand <= n_ctx) { + return runnable; + } + } + bool draining = false; + for (const auto & slot : slots) { + draining |= slot.task && slot.task->id == preempt_drain_task && slot.state != SLOT_STATE_PREEMPTED; + } + if (!draining) { + preempt_drain_task = -1; + std::vector waiting; + for (auto & slot : slots) { + if (slot.state == SLOT_STATE_PREEMPTED && !slot.preempt_manual) { + waiting.push_back(&slot); + } + } + if (!waiting.empty()) { + // A completed survivor can leave interleaved cells throughout the + // pool. Clear idle caches before restoring, even if the cell count + // already fits, to avoid thousands of small scatter transfers. + while (try_clear_idle_slots()) {} + } + std::sort(waiting.begin(), waiting.end(), [](const server_slot * a, const server_slot * b) { + if (a->effective_priority() != b->effective_priority()) { + return a->effective_priority() > b->effective_priority(); + } + if (a->n_preempt != b->n_preempt) { + return a->n_preempt > b->n_preempt; + } + return a->preempt_since < b->preempt_since; + }); + for (auto * slot : waiting) { + const int64_t needed = slot->prompt.tokens.pos_next() + preempt_step(*slot); + while (preempt_demand() + needed > n_ctx && try_clear_idle_slots()) {} + if (preempt_demand() + needed <= n_ctx) { + preempt_restore(*slot); // fit-first: skip a larger snapshot that does not fit + } + } + } + + // Test-only round trips happen at the same safe boundary as real parks. + // The interval is crossed rather than tested for equality: MTP accepts + // several tokens in one step. Record before parking to avoid retriggering. + if (preempt_every > 0) { + for (auto & slot : slots) { + if (slot.state == SLOT_STATE_GENERATING && preempt_candidate(slot) && + slot.stats.n_gen / preempt_every > slot.preempt_forced_at) { + slot.preempt_forced_at = slot.stats.n_gen / preempt_every; + if (preempt_park(slot, false)) { + preempt_restore(slot); + } + } + } + } + + while (preempt_demand() > n_ctx && try_clear_idle_slots()) {} + if (preempt_demand() > n_ctx || draining) { + std::vector active; + for (auto & slot : slots) { + if (preempt_candidate(slot)) { + active.push_back(&slot); + } + } + // Candidates are the lowest-priority tier. Protect the longest + // sequence within that tier unless it is the only candidate; a + // lower-priority long request must not displace an important one. + // Newest task wins ties among otherwise eligible victims. + std::sort(active.begin(), active.end(), [](const server_slot * a, const server_slot * b) { + if (a->effective_priority() != b->effective_priority()) { + return a->effective_priority() < b->effective_priority(); + } + return a->task->id > b->task->id; + }); + size_t remaining = active.size(); + std::vector candidates = active; + while (remaining > 1 && !candidates.empty()) { + server_slot * longest = nullptr; + for (auto * slot : active) { + if (slot->state != SLOT_STATE_PREEMPTED && + (!longest || slot->prompt.n_tokens() > longest->prompt.n_tokens() || + (slot->prompt.n_tokens() == longest->prompt.n_tokens() && + slot->effective_priority() > longest->effective_priority()))) { + longest = slot; + } + } + auto it = candidates.begin(); + if ((*it)->task->id == preempt_drain_task) { + candidates.erase(it); + continue; + } + if (*it == longest && !draining) { + auto alternative = std::find_if(it + 1, candidates.end(), [&](const server_slot * other) { + return other->effective_priority() == longest->effective_priority() && + other->task->id != preempt_drain_task; + }); + if (alternative != candidates.end()) { + it = alternative; + } + } + auto * victim = *it; + candidates.erase(it); + if (preempt_park(*victim, false)) { + --remaining; + } + } + if (remaining == 1) { + for (auto * slot : active) { + if (slot->state != SLOT_STATE_PREEMPTED) { + preempt_drain_task = slot->task->id; + } + } + } + if (preempt_demand() > n_ctx) { + // RAM exhaustion or an unparkable shared-prompt group: preserve + // live tasks. A cancellation/control request can release room. + ++metrics.n_preempt_blocked; + return false; + } + } + for (const auto & slot : slots) { + if (slot.is_processing() && slot.state != SLOT_STATE_PREEMPTED) { + return true; + } + } + return false; + } + void update_slots() { #ifdef DEBUG_TIMINGS static int64_t t_prev = 0; @@ -2688,6 +3029,13 @@ struct server_context_impl { } #endif + if (preempt_enabled() && !preempt_schedule()) { + // No runnable slots (e.g. all explicitly parked). The task queue sleeps + // until a control request or cancellation arrives; do not busy-loop. + metrics_flush_idle(); + return; + } + // check if all slots are idle { bool all_idle = true; @@ -3000,7 +3348,7 @@ struct server_context_impl { return; // batch is full, skip remaining slots } - if (!slot.is_processing()) { + if (!slot.is_processing() || slot.state == SLOT_STATE_PREEMPTED) { return; } @@ -3404,8 +3752,14 @@ struct server_context_impl { const auto & spans = slot.task->params.message_spans; const auto last_user_pos = spans.last_user_message_pos(); + int32_t prompt_batch_limit = n_batch; + if (preempt_enabled()) { + prompt_batch_limit = std::min(n_batch, + batch.size() + std::max(0, n_ctx - preempt_resident())); + } + // add prompt tokens for processing in the current batch - while (slot.prompt.n_tokens() < slot.task->n_tokens() && batch.size() < n_batch) { + while (slot.prompt.n_tokens() < slot.task->n_tokens() && batch.size() < prompt_batch_limit) { // get next token to process llama_token cur_tok = input_tokens[slot.prompt.n_tokens()]; if (cur_tok == LLAMA_TOKEN_NULL) { @@ -4642,11 +4996,6 @@ void server_routes::init_routes() { this->post_slots = [this](const server_http_req & req) { auto res = create_response(); - if (params.slot_save_path.empty()) { - res->error(format_error_response("This server does not support slots action. Start it with `--slot-save-path`", ERROR_TYPE_NOT_SUPPORTED)); - return res; - } - std::string id_slot_str = req.get_param("id_slot"); int id_slot; @@ -4659,6 +5008,30 @@ void server_routes::init_routes() { std::string action = req.get_param("action"); + if (action == "park" || action == "unpark") { + if (!params.endpoint_slots) { + res->error(format_error_response("Slot control requires --slots", ERROR_TYPE_NOT_SUPPORTED)); + return res; + } + server_task task(action == "park" ? SERVER_TASK_TYPE_SLOT_PARK : SERVER_TASK_TYPE_SLOT_UNPARK); + task.id = res->rd.get_new_id(); + task.slot_action.id_slot = id_slot; + res->rd.post_task(std::move(task), true); + auto result = res->rd.next(req.should_stop); + if (!result) { + return res; + } + if (result->is_error()) { + res->error(result->to_json()); + } else { + res->ok(result->to_json()); + } + return res; + } + if (params.slot_save_path.empty()) { + res->error(format_error_response("This server does not support slots action. Start it with `--slot-save-path`", ERROR_TYPE_NOT_SUPPORTED)); + return res; + } if (action == "save") { return handle_slots_save(req, id_slot); } diff --git a/tools/server/server-queue.cpp b/tools/server/server-queue.cpp index 78169e9a5d8..f7fdb072995 100644 --- a/tools/server/server-queue.cpp +++ b/tools/server/server-queue.cpp @@ -448,6 +448,9 @@ server_task_result_ptr server_response::recv(const std::unordered_set & id_ } server_task_result_ptr server_response::recv_with_timeout(const std::unordered_set & id_tasks, int timeout) { + // Other streams and /slots polling also notify this condition. They must + // not keep extending a parked stream's heartbeat/cancellation timeout. + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(timeout); while (true) { std::unique_lock lock(mutex_results); @@ -459,7 +462,7 @@ server_task_result_ptr server_response::recv_with_timeout(const std::unordered_s } } - std::cv_status cr_res = condition_results.wait_for(lock, std::chrono::seconds(timeout)); + std::cv_status cr_res = condition_results.wait_until(lock, deadline); if (!running) { RES_DBG("%s : queue result stop\n", __func__); std::terminate(); // we cannot return here since the caller is HTTP code diff --git a/tools/server/server-schema.cpp b/tools/server/server-schema.cpp index 64b9251295c..0c230632d4b 100644 --- a/tools/server/server-schema.cpp +++ b/tools/server/server-schema.cpp @@ -14,6 +14,10 @@ std::vector> make_llama_cmpl_schema(const common_params & fields.emplace_back(f); }; + add((new field_num("priority", params.priority)) + ->set_hard_limits(INT32_MIN, INT32_MAX) + ->set_desc("Scheduling priority, higher values are more important (default: 0)")); + add((new field_bool("verbose", params.verbose)) ->set_desc("Include __verbose field in the response with additional debug information")); @@ -518,6 +522,10 @@ task_params eval_llama_cmpl_schema( const std::vector & logit_bias_eog, const json & data) { task_params params; + if (data.contains("priority") && (!data.at("priority").is_number_integer() || + data.at("priority").get() < INT32_MIN || data.at("priority").get() > INT32_MAX)) { + throw std::invalid_argument("priority must be a signed 32-bit integer"); + } // Sampling parameter defaults are loaded from the global server context (but individual requests can still override them) params.sampling = params_base.sampling; diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index 0d3beb313ce..ecc2c3d2009 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -41,6 +41,7 @@ json task_params::to_json(bool only_metrics) const { if (only_metrics) { return json { + {"priority", priority}, {"seed", sampling.seed}, {"temperature", sampling.temp}, {"dynatemp_range", sampling.dynatemp_range}, @@ -93,6 +94,7 @@ json task_params::to_json(bool only_metrics) const { } return json { + {"priority", priority}, {"seed", sampling.seed}, {"temperature", sampling.temp}, {"dynatemp_range", sampling.dynatemp_range}, @@ -1522,6 +1524,10 @@ json server_task_result_metrics::to_json() { // metrics definition: https://prometheus.io/docs/practices/naming/#metric-names std::string server_task_result_metrics::to_metrics() { const std::vector counters = { + { "preemptions_total", "Sequences parked", (double) metrics.n_preempt }, + { "preempt_restores_total", "Sequences restored", (double) metrics.n_restore }, + { "preempt_blocked_total", "Scheduling attempts blocked by snapshot RAM or policy", (double) metrics.n_preempt_blocked }, + { "preempt_copy_seconds_total", "Time copying sequence snapshots to and from host RAM", metrics.preempt_copy_us / 1.e6 }, { "prompt_tokens_total", "Number of prompt tokens processed, excluding cached tokens", @@ -1566,6 +1572,8 @@ std::string server_task_result_metrics::to_metrics() { }; const std::vector gauges = { + { "preempt_ram_bytes", "Host RAM held by sequence snapshots", (double) metrics.preempt_ram_bytes }, + { "requests_preempted", "Requests parked with live streams", (double) metrics.n_preempted }, { "prompt_tokens_seconds", "Average prompt throughput in tokens/s", diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 9c99143f8e1..d9ed469a780 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -25,6 +25,8 @@ enum server_task_type { SERVER_TASK_TYPE_SLOT_SAVE, SERVER_TASK_TYPE_SLOT_RESTORE, SERVER_TASK_TYPE_SLOT_ERASE, + SERVER_TASK_TYPE_SLOT_PARK, + SERVER_TASK_TYPE_SLOT_UNPARK, SERVER_TASK_TYPE_GET_LORA, SERVER_TASK_TYPE_SET_LORA, }; @@ -48,6 +50,7 @@ enum stop_type { }; struct task_params { + int32_t priority = 0; // higher values are more important bool stream = false; bool include_usage = false; bool cache_prompt = true; // remember the prompt to avoid reprocessing all prompt diff --git a/tools/server/tests/unit/test_preempt_priority.py b/tools/server/tests/unit/test_preempt_priority.py new file mode 100644 index 00000000000..59b896ce253 --- /dev/null +++ b/tools/server/tests/unit/test_preempt_priority.py @@ -0,0 +1,198 @@ +"""Request preemption tests. Use a shared two-slot pool and actual SSE streams.""" +import json +import time +from concurrent.futures import ThreadPoolExecutor + +import pytest +import requests + +from utils import ServerPreset + + +def start(monkeypatch, ctx=1024, every=0): + monkeypatch.setenv("LLAMA_SERVER_PREEMPT_EVERY", str(every)) + server = ServerPreset.tinyllama2() + server.n_slots = 2 + server.n_ctx = ctx + server.n_batch = 64 + server.n_ubatch = 64 + server.n_threads = 1 + server.n_gpu_layer = 0 # CPU reference isolates state preservation from CUDA batch/layout rounding. + server.n_predict = -1 + server.kv_unified = True + server.server_continuous_batching = True + server.server_slots = True + server.server_metrics = True + server.start() + return server + + +def url(server, path): + return f"http://{server.server_host}:{server.server_port}{path}" + + +def body(slot=0, n=600, priority=0, prompt="Once upon a time"): + return dict(prompt=prompt, id_slot=slot, n_predict=n, temperature=0, + seed=123, ignore_eos=True, cache_prompt=False, + return_tokens=True, priority=priority) + + +def metrics(server): + response = requests.get(url(server, "/metrics"), timeout=10) + response.raise_for_status() + return {line.split()[0]: float(line.split()[1]) + for line in response.text.splitlines() + if line.startswith("llamacpp:") and "{" not in line} + + +def slots(server): + response = requests.get(url(server, "/slots"), timeout=10) + response.raise_for_status() + return response.json() + + +def stream(server, payload, chunks): + with requests.post(url(server, "/completion"), json={**payload, "stream": True}, + stream=True, timeout=60) as response: + response.raise_for_status() + for line in response.iter_lines(chunk_size=1): + if line.startswith(b"data: "): + item = json.loads(line[6:]) + assert "error" not in item, item + chunks.append(item) + assert chunks[-1]["stop"] + return "".join(item.get("content", "") for item in chunks) + + +def until(predicate, timeout=20): + end = time.monotonic() + timeout + while time.monotonic() < end: + value = predicate() + if value: + return value + time.sleep(0.005) + raise AssertionError("condition not reached before timeout") + + +def generated(chunks): + return sum(len(item.get("tokens", [])) for item in chunks if not item.get("stop")) + + +@pytest.mark.parametrize("prompt", ["Once upon a time", "The little dog went"]) +def test_forced_identity(monkeypatch, prompt): + outputs = [] + for every in (0, 30): + server = start(monkeypatch, every=every) + response = requests.post(url(server, "/completion"), json=body(n=240, prompt=prompt), timeout=60) + response.raise_for_status() + outputs.append(response.json()) + assert metrics(server)["llamacpp:preemptions_total"] == (7 if every else 0) + assert metrics(server)["llamacpp:preempt_ram_bytes"] == 0 + server.stop() + assert outputs[0]["tokens"] == outputs[1]["tokens"] + assert outputs[0]["content"].encode() == outputs[1]["content"].encode() + + +def test_pool_pressure_and_priority(monkeypatch): + server = start(monkeypatch) + chunks = [[], []] + observed = [] + with ThreadPoolExecutor(2) as pool: + # Give the low-priority request a lead so it is the longest. It must + # still yield when it is the sole member of the lowest-priority tier. + low = pool.submit(stream, server, body(1, priority=0), chunks[1]) + until(lambda: generated(chunks[1]) >= 30) + high = pool.submit(stream, server, body(0, priority=10), chunks[0]) + while not (high.done() and low.done()): + observed.extend(slots(server)) + time.sleep(0.005) + high.result() + low.result() + # Each request is < 1024 tokens; together they exceed the shared pool. + assert all(generated(items) == 600 for items in chunks) + parked = [s for s in observed if s["is_preempted"]] + assert parked + assert all(s["priority"] == 0 and s["id"] == 1 for s in parked) + assert metrics(server)["llamacpp:preemptions_total"] > 0 + assert metrics(server)["llamacpp:requests_preempted"] == 0 + + +def test_explicit_round_trip(monkeypatch): + server = start(monkeypatch, ctx=4096) + payload = body(n=1800) + reference = requests.post(url(server, "/completion"), json=payload, timeout=60).json() + chunks = [] + with ThreadPoolExecutor(1) as pool: + future = pool.submit(stream, server, payload, chunks) + until(lambda: generated(chunks) >= 100) + response = requests.post(url(server, "/slots/0?action=park"), timeout=10) + response.raise_for_status() + parked = response.json() + assert parked["is_preempted"] and parked["preempt_ram_bytes"] > 0 + assert parked["preempt_manual"] + # The client can drain already-sent chunks, but the server generates no more. + before = slots(server)[0]["next_token"][0]["n_decoded"] + other = requests.post(url(server, "/completion"), json=body(1, n=100), timeout=60) + other.raise_for_status() + assert other.json()["tokens_predicted"] == 100 + after = slots(server)[0] + assert after["is_preempted"] + assert after["next_token"][0]["n_decoded"] == before + assert metrics(server)["llamacpp:requests_preempted"] == 1 + response = requests.post(url(server, "/slots/0?action=unpark"), timeout=10) + response.raise_for_status() + assert not response.json()["preempt_manual"] + content = future.result(timeout=60) + assert content.encode() == reference["content"].encode() + tokens = [token for item in chunks if not item.get("stop") for token in item.get("tokens", [])] + assert tokens == reference["tokens"] + assert metrics(server)["llamacpp:preempt_ram_bytes"] == 0 + + +def test_metrics_exposed(monkeypatch): + server = start(monkeypatch) + text = requests.get(url(server, "/metrics"), timeout=10).text + for name, kind in [("preemptions_total", "counter"), ("preempt_restores_total", "counter"), + ("preempt_blocked_total", "counter"), ("preempt_copy_seconds_total", "counter"), + ("preempt_ram_bytes", "gauge"), ("requests_preempted", "gauge")]: + assert f"# TYPE llamacpp:{name} {kind}" in text + for slot in slots(server): + assert slot["priority"] == 0 and not slot["is_preempted"] + assert slot["n_preempt"] == 0 and slot["preempt_ram_bytes"] == 0 + + +@pytest.mark.parametrize("priority", [1.5, "10", True, 2**31, -(2**31)-1]) +def test_priority_validation(monkeypatch, priority): + server = start(monkeypatch) + response = requests.post(url(server, "/completion"), json=body(n=1, priority=priority), timeout=10) + assert response.status_code == 400 + + +def test_cancel_parked_releases_ram(monkeypatch): + server = start(monkeypatch, ctx=4096) + response = requests.post(url(server, "/completion"), json={**body(n=1800), "stream": True, "sse_ping_interval": 1}, + headers={"X-Conversation-Id": "preempt-cancel"}, + stream=True, timeout=60) + try: + until(lambda: slots(server)[0].get("next_token", [{}])[0].get("n_decoded", 0) > 10) + parked = requests.post(url(server, "/slots/0?action=park"), timeout=10) + parked.raise_for_status() + assert parked.json()["preempt_ram_bytes"] > 0 + finally: + requests.delete(url(server, "/v1/stream?conv_id=preempt-cancel"), timeout=10).raise_for_status() + response.close() + until(lambda: not slots(server)[0]["is_processing"]) + assert metrics(server)["llamacpp:preempt_ram_bytes"] == 0 + assert slots(server)[0]["n_prompt_tokens"] == 0 + + +def test_snapshot_ram_bound(monkeypatch): + monkeypatch.setenv("LLAMA_ARG_PREEMPT_RAM", "1") + server = start(monkeypatch, ctx=2048, every=1700) + # A 1700-token stories260K snapshot needs > 1 MiB. Refusing the + # diagnostic round trip must not damage the still-resident request. + response = requests.post(url(server, "/completion"), json=body(n=1800), timeout=60) + response.raise_for_status() + assert response.json()["tokens_predicted"] == 1800 + assert metrics(server)["llamacpp:preemptions_total"] == 0 + assert metrics(server)["llamacpp:preempt_ram_bytes"] == 0 diff --git a/tools/server/tests/utils.py b/tools/server/tests/utils.py index a0d2dfa3c59..57e9d6c43cf 100644 --- a/tools/server/tests/utils.py +++ b/tools/server/tests/utils.py @@ -188,7 +188,7 @@ def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None: server_args.extend(["--ubatch-size", self.n_ubatch]) if self.n_threads: server_args.extend(["--threads", self.n_threads]) - if self.n_gpu_layer: + if self.n_gpu_layer is not None: server_args.extend(["--n-gpu-layers", self.n_gpu_layer]) if self.server_continuous_batching: server_args.append("--cont-batching")