diff --git a/atom/model_engine/block_manager.py b/atom/model_engine/block_manager.py index 3e96ab7ccf..caa9504d47 100644 --- a/atom/model_engine/block_manager.py +++ b/atom/model_engine/block_manager.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from math import inf, isinf from time import monotonic +from weakref import WeakKeyDictionary import numpy as np import xxhash @@ -131,6 +132,11 @@ def __init__( # tokens (see _hash_block_size). == block_size when DCP is off. self.hash_block_size = self.block_size * self.dcp_world_size self.enable_prefix_caching = config.enable_prefix_caching + # Content hashes only: pool hits and resource fit are always rechecked. + # Weak keys keep this scheduler-side cache out of request serialization. + self._prefill_probe_hashes: WeakKeyDictionary[ + Sequence, tuple[int, list[int]] + ] = WeakKeyDictionary() self.total_evicted_blocks: int = 0 kv_events = getattr(config, "kv_events_config", None) @@ -861,17 +867,22 @@ def _chain_to( chain.append(h) return chain - def can_allocate(self, seq: Sequence, record: bool = True) -> int: + def can_allocate( + self, + seq: Sequence, + record: bool = True, + *, + block_hashes: list[int] | None = None, + reuse_hashes: bool = False, + ) -> int: """Return number of cache-hit blocks (>=0) if seq fits, else -1. - `record=False` marks a fit probe -- asking only whether the seq *could* - be admitted. The fit answer is identical, but the probe is not - side-effect-free: the instrumentation and checkpoint-demand/-end writes - below run before the `record` gate. That is safe -- the sole probe caller - reads only the `>= 0` return, and the next real `can_allocate` overwrites - those fields first. `record` gates only the joint-boundary commit (seq - joint fields + funnel counters), so a probe cannot inflate the - operator-visible funnel (see `_commit_joint_boundary`). + `record=False` returns the same fit and HBM hit without committing + joint-load fields or joint-boundary counters. Checkpoint demand/end + and hit instrumentation still refresh; demand counters deduplicate per + request. An optional `block_hashes` output + lets the scheduler defer `record_allocation` until after its wait and + token-budget checks, reusing this probe's chain without another walk. The hit count is the contiguous run of cache hits starting at the prompt's first block. On the first miss we break: subsequent blocks @@ -888,6 +899,10 @@ def can_allocate(self, seq: Sequence, record: bool = True) -> int: # The full per-request width, because that is what `allocate` will take: # gating on one slot would admit a request the pool cannot give a # rollback set to. + if block_hashes is None: + block_hashes = [] + else: + block_hashes.clear() if seq.has_per_req_cache and not self.state.has_free(self.state_slots_per_req): return -1 if not self.enable_prefix_caching: @@ -899,10 +914,21 @@ def can_allocate(self, seq: Sequence, record: bool = True) -> int: # match). Record each block's hash for the SWA scan below. h = seq.cache_seed compressed_hit = 0 - block_hashes: list[int] = [] + cached_hashes = None + if reuse_hashes: + seed, cached_hashes = self._prefill_probe_hashes.get(seq, (h, [])) + if seed != h: + cached_hashes = [] + self._prefill_probe_hashes[seq] = (h, cached_hashes) + immutable_blocks = seq.num_prompt_tokens // self.hash_block_size for i in range(self._n_hash_blocks(seq) - 1): token_ids = self._hash_block_tokens(seq, i) - h = self.compute_hash(token_ids, h) + if cached_hashes is not None and i < len(cached_hashes): + h = cached_hashes[i] + else: + h = self.compute_hash(token_ids, h) + if cached_hashes is not None and i < immutable_blocks: + cached_hashes.append(h) block_id = self.kv.lookup(h) if block_id == -1 or self.kv.block(block_id).token_ids != token_ids: break @@ -949,25 +975,11 @@ def can_allocate(self, seq: Sequence, record: bool = True) -> int: self._record_checkpoint_end(seq) if not self._has_page_units(num_new_blocks, protected_hash): return -1 - # A boundary LMCache and the tier can jointly reach, above this hit. - # Committed to the seq rather than returned: what `allocate` claims from - # HBM is still `num_cached_blocks`, and the joint boundary only decides - # where the two loads are aimed. `record` is the probe/admission split: - # a real admission computes the boundary and commits it (seq fields + - # funnel counters); a fit probe skips it entirely. The `num_cached_blocks` - # returned below does not depend on the decision, so gating the - # computation -- not just the commit -- spares a probe the O(prompt) work - # for a value it would only discard: `_joint_kv_boundary` walks a chained - # xxhash up to the LMCache-only cap (`_chain_to`) plus a `_gated_hit` - # rescan. A 128k prompt at the front of a KV-pressured queue paid that - # full chain on every scheduling pass (`is_mixed_batch` peeks up to four - # waiting seqs with `record=False`) for a decision nothing read. Placed - # below the refusal, not above it, for the same reason `_extend_hash_chain` - # is: a refused admission would discard it, and nothing between here and - # the refusal reads the seq's joint fields -- this is that move's twin. + # A direct admission commits its joint boundary here. The scheduler + # probes first and calls record_allocation only after dependency and + # budget checks; refused probes avoid the extra LMCache hash walk. if record: - decision = self._joint_kv_boundary(seq, num_cached_blocks, block_hashes) - self._commit_joint_boundary(seq, decision) + self.record_allocation(seq, num_cached_blocks, block_hashes) # After the refusal, not before it. The chain is O(prompt) xxhash plus # two temporaries per block, and a refused admission discards it — a # 128k prompt queued behind a full pool paid ~2000 rounds per waiting @@ -983,6 +995,14 @@ def can_allocate(self, seq: Sequence, record: bool = True) -> int: self._extend_hash_chain(seq, block_hashes) return num_cached_blocks + def record_allocation( + self, seq: Sequence, num_cached_blocks: int, block_hashes: list[int] + ) -> None: + """Commit the fit probe before any allocation changes the pool.""" + if self.enable_prefix_caching: + decision = self._joint_kv_boundary(seq, num_cached_blocks, block_hashes) + self._commit_joint_boundary(seq, decision) + def allocate(self, seq: Sequence, num_cached_blocks: int = 0) -> bool: """Allocate blocks for `seq`. `num_cached_blocks` is the hit count returned by `can_allocate` (0 if caller didn't call it). @@ -2060,7 +2080,9 @@ def cancel_midstep(self, seq: Sequence) -> None: self.state.cancel_midstep(seq.midstep_reservations) seq.midstep_reservations = [] - def checkpoint_cut(self, seq: Sequence, start: int, end: int) -> int: + def checkpoint_cut( + self, seq: Sequence, start: int, end: int, *, record: bool = True + ) -> int: """Earliest ladder position in `(start, end]`, or 0 if there is none. What a prefill chunk is cut at so its forward lands exactly on a rung. @@ -2133,7 +2155,7 @@ class still needs the forward to land there, and the readable ones lose # `chunks_cut_for_demand` would swamp the convergence signal that # counter exists to expose. The demand is checked first because when # the two coincide it is the demand that evidenced the position. - if target != rung and target < end: + if record and target != rung and target < end: if target == demand: self.chunks_cut_for_demand += 1 else: @@ -2607,6 +2629,7 @@ def deallocate(self, seq: Sequence): seq.state_fork_src = -1 seq.offload_joint.load_hash = -1 + seq.offload_joint.reset_joint() def can_append(self, seq: Sequence, num_new_tokens: int = 1) -> bool: seq_len = len(seq) diff --git a/atom/model_engine/engine_core.py b/atom/model_engine/engine_core.py index 320415a763..319ab8edce 100644 --- a/atom/model_engine/engine_core.py +++ b/atom/model_engine/engine_core.py @@ -167,6 +167,12 @@ def __init__(self, config: Config, input_address: str, output_address: str): config, state_runtime=self.state_runtime, ) + if ( + config.parallel_config.data_parallel_size == 1 + and config.pipeline_parallel_size == 1 + and envs.ATOM_PREFILL_DECODE_INTERVAL > 0 + ): + self._init_prefill_delayer(config) self.kv_transfer_enabled = bool(config.kv_transfer_config) self._next_idle_kv_drain = 0.0 @@ -188,6 +194,31 @@ def __init__(self, config: Config, input_address: str, output_address: str): self._send_ready_signal() logger.info(f"{self.label}: EngineCore fully initialized and ready") + def _init_prefill_delayer(self, config: Config, cpu_group=None): + if ( + not envs.ATOM_ENABLE_PREFILL_DELAYER + or config.enable_rapidserve + or self.scheduler is None + ): + return + from atom.model_engine.prefill_delayer import PrefillDelayer + + self.scheduler.set_prefill_delayer( + PrefillDelayer( + dp_size=config.parallel_config.data_parallel_size, + cpu_group=cpu_group, + max_num_batched_tokens=config.max_num_batched_tokens, + target_fill=envs.ATOM_PREFILL_DELAYER_TARGET_FILL, + ttft_max_ticks=envs.ATOM_PREFILL_DELAYER_TTFT_MAX_TICKS, + partial_max_ticks=envs.ATOM_PREFILL_DELAYER_PARTIAL_MAX_TICKS, + stall_ticks=envs.ATOM_PREFILL_DELAYER_STALL_TICKS, + kv_high_watermark=envs.ATOM_PREFILL_DELAYER_KV_HIGH_WATERMARK, + token_usage_low_watermark=envs.ATOM_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK, + max_queue_ms=envs.ATOM_PREFILL_DELAYER_MAX_QUEUE_MS, + prefill_decode_interval=envs.ATOM_PREFILL_DECODE_INTERVAL, + ) + ) + def _freeze_after_startup(self): """Freeze this process and its ModelRunner workers. @@ -726,24 +757,7 @@ def __init__(self, config: Config, input_address: str, output_address: str): self.engines_running = True self._shutting_down = False - if envs.ATOM_ENABLE_PREFILL_DELAYER: - from atom.model_engine.prefill_delayer import PrefillDelayer - - self.scheduler.set_prefill_delayer( - PrefillDelayer( - dp_size=config.parallel_config.data_parallel_size, - cpu_group=self.dp_group, - max_num_batched_tokens=config.max_num_batched_tokens, - target_fill=envs.ATOM_PREFILL_DELAYER_TARGET_FILL, - ttft_max_ticks=envs.ATOM_PREFILL_DELAYER_TTFT_MAX_TICKS, - partial_max_ticks=envs.ATOM_PREFILL_DELAYER_PARTIAL_MAX_TICKS, - stall_ticks=envs.ATOM_PREFILL_DELAYER_STALL_TICKS, - kv_high_watermark=envs.ATOM_PREFILL_DELAYER_KV_HIGH_WATERMARK, - token_usage_low_watermark=envs.ATOM_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK, - max_queue_ms=envs.ATOM_PREFILL_DELAYER_MAX_QUEUE_MS, - prefill_decode_interval=envs.ATOM_PREFILL_DECODE_INTERVAL, - ) - ) + self._init_prefill_delayer(config, self.dp_group) def _init_data_parallel(self, config: Config): dp_rank = config.parallel_config.data_parallel_rank diff --git a/atom/model_engine/engine_utility.py b/atom/model_engine/engine_utility.py index 808c3e0d99..1e131bdc27 100644 --- a/atom/model_engine/engine_utility.py +++ b/atom/model_engine/engine_utility.py @@ -218,16 +218,19 @@ def _handle_clear_kv_cache(self, args: dict): ) def _handle_abort_request(self, args: dict): - """Mark a sequence ABORTED (client disconnected) so the scheduler finishes - it at the next step via the normal stop path (frees KV, drops it).""" + """Cancel queued work promptly; running forwards finish via postprocess.""" req_id = args.get("req_id") if isinstance(args, dict) else None if req_id is None or self.scheduler is None: return - found = False - for seq in list(self.scheduler.running) + list(self.scheduler.waiting): - if seq.id == req_id: - seq.status = SequenceStatus.ABORTED - found = True + abort = getattr(self.scheduler, "abort_request", None) + if callable(abort): + found = abort(req_id) + else: + found = False + for seq in list(self.scheduler.running) + list(self.scheduler.waiting): + if seq.id == req_id: + seq.status = SequenceStatus.ABORTED + found = True logger.info(f"{self.label}: abort_request req_id={req_id} found={found}") def _handle_configure_hidden_states(self, args: dict): diff --git a/atom/model_engine/prefill_delayer.py b/atom/model_engine/prefill_delayer.py index fdebd50255..ca6c20776c 100644 --- a/atom/model_engine/prefill_delayer.py +++ b/atom/model_engine/prefill_delayer.py @@ -1,5 +1,5 @@ """ -PrefillDelayer — a cross-DP-rank prefill *coalescer* for ATOM. +PrefillDelayer — local TP or cross-DP prefill coalescing for ATOM. Purpose ------- @@ -12,16 +12,21 @@ The delayer has two related jobs: **hold back prefill admission until the accumulated prefill is worth a forward, then release** — Nagle's algorithm for prefill — and optionally protect a fixed number of decode passes after every -prefill. While it holds, decode keeps running; TTFT is bounded so a held request -never starves. +prefill. While it holds, decode keeps running. Coalescing hold episodes are +bounded; the decode interval, backlog, and execution add to observed TTFT. Single-rank / TP-only mode -------------------------- -With ``cpu_group=None`` (``dp_size=1``) the delayer runs the same coalescer +With ``is_local`` (``dp_size=1`` and ``cpu_group=None``) the delayer runs the same coalescer locally and skips the cross-rank ``all_reduce`` — a single scheduler drives all TP workers, so there is no cross-rank phase to align. Only the batching value remains (fill / stall / ttft / kv bounds decide FIRE/HOLD from this rank's own counts); the ``n_prefillable < dp_size`` alignment gate is vacuous. +EngineCore opts TP/DCP into this mode when DP=1, PP=1, the master switch is +on, and ``ATOM_PREFILL_DECODE_INTERVAL > 0``. During local decode protection, +``protects_decode`` lets the scheduler skip cache probes and queue-age scans. +The interval precedes all coalescing bounds, including ``max_queue_ms``; those +bounds do not impose a hard end-to-end TTFT limit. It is NOT about mixing prefill+decode in one forward. It DOES preserve cross-DP phase alignment: it only releases when every rank is prefill-ready (so all ranks @@ -47,6 +52,7 @@ # -- must-fire bounds (release even if unaligned / underfilled) -- if G_running_dec == 0: FIRE # no decode to hide the wait behind if any_kv_high or any_kv_low: FIRE # KV pressure / starvation + if any_queue_hot: FIRE # queue age bound if hold_ticks >= ttft_max_ticks: FIRE # TTFT bound if any_partial and hold_ticks >= partial_max_ticks: FIRE # partial holds KV — tight bound # -- alignment gate: never fire while some rank lacks prefill (anti-skew) -- @@ -138,6 +144,8 @@ def __init__( max_queue_ms: float | None = None, prefill_decode_interval: int = 0, ): + if dp_size > 1 and cpu_group is None: + raise ValueError("Cross-DP prefill coalescing requires a CPU group") self.dp_size = dp_size self.cpu_group = cpu_group self.max_num_batched_tokens = max_num_batched_tokens @@ -250,6 +258,25 @@ def _clamp_ticks(name: str, value: int) -> int: return 1 return value + @property + def is_local(self) -> bool: + return self.dp_size == 1 and self.cpu_group is None + + def protects_decode(self, running_decode_batch: int) -> bool: + """Whether the local decision can skip cache probes during protection.""" + return ( + self.is_local + and not self._first + and running_decode_batch > 0 + and ( + self._decode_interval_remaining > 0 + or ( + self._prefill_executed_since_last_decision + and self.prefill_decode_interval > 0 + ) + ) + ) + def should_allow_prefill( self, prefillable: bool, @@ -267,11 +294,14 @@ def should_allow_prefill( Args: prefillable: this rank has admittable prefill work (fresh head that can allocate, or a resumable partial). Only prefillable ranks - count toward the fill target and the alignment gate. + count toward the fill target and the alignment gate. During + local decode protection, the caller uses work existence alone: + fit cannot change the hard interval decision. pending_tokens: this rank's accumulated prefill tokens — fresh waiting new-tokens PLUS the remaining tokens of resumable partials — already capped at max_num_batched_tokens by the - caller. The coalescer's fill signal. + caller. The local scheduler bounds its scan and uses chunk + limits, so this signal can omit work beyond that scan. running_decode_batch: decode seqs running on this rank (NOT counting mid-chunked-prefill seqs). If no rank has decode, holding wastes GPU → fire. @@ -279,11 +309,12 @@ def should_allow_prefill( KV-high (can't accumulate more) and KV-low (GPU starving) bounds. has_partial: this rank has a mid-chunked-prefill seq in flight. Its remaining tokens are in pending_tokens; a partial only forces - release once held for partial_max_ticks (it holds KV). + release once held for partial_max_ticks (it holds KV), counted + after the decode-protection interval. oldest_waiting_age_ms: age (ms since arrival) of this rank's oldest schedulable waiting prefill. If it reaches max_queue_ms, this - rank flags the TTFT SLA guard and all ranks release. 0 / no - waiting prefill / max_queue_ms=None → guard inactive. + rank flags the age guard; all ranks release after decode + protection. No waiting prefill / max_queue_ms=None → inactive. """ # KV + queue-age bounds are gated on this rank actually having prefill to # push (firing when this rank can't admit anything would be a no-op). @@ -353,8 +384,8 @@ def should_allow_prefill( return True # Hard post-prefill decode protection. The countdown is identical on - # all ranks because it is armed from the globally reduced prefillable - # count and schedule() calls this method in lockstep. Do not advance the + # all ranks because it is armed from the globally reduced execution + # signal and schedule() calls this method in lockstep. Do not advance the # coalescer's own hold/TTFT episode: SGLang applies the interval before # invoking PrefillDelayer, so its delay budget starts afterwards. if self._decode_interval_remaining > 0: @@ -377,10 +408,8 @@ def should_allow_prefill( return self._fire("nodecode", g_pending) if any_kv_high or any_kv_low: return self._fire("kv", g_pending) - # TTFT SLA guard: a real request has queued (since arrival) past the - # threshold — release now regardless of fill/alignment. This is the - # end-to-end wait (includes backlog + coalescer holds), unlike the - # tick-based ttft bound below which only caps a single hold episode. + # Queue age includes backlog and prior holds. This release applies + # after decode protection; it is not an end-to-end TTFT guarantee. if any_queue_hot: return self._fire("queue_ms", g_pending) if self._hold_ticks >= self.ttft_max_ticks: diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index 9bc823a3bb..9e1233acb0 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -26,6 +26,7 @@ import time from collections import deque from collections.abc import Iterable +from typing import TYPE_CHECKING import numpy as np @@ -49,6 +50,9 @@ ) from atom.utils import envs +if TYPE_CHECKING: + from atom.model_engine.prefill_delayer import PrefillDelayer + logger = logging.getLogger("atom") @@ -698,14 +702,172 @@ def __init__( endpoint="", ) - # Cross-DP prefill alignment. Set by DPEngineCoreProc after - # dp_group is available. See `prefill_delayer.py` for rationale. - from atom.model_engine.prefill_delayer import PrefillDelayer - + # Set by EngineCore for cross-DP alignment or opt-in TP decode protection. self.prefill_delayer: PrefillDelayer | None = None + self._local_prefill_coalescing = False + self._inflight_prefix_wait: dict[int, int] = {} - def set_prefill_delayer(self, delayer) -> None: + def set_prefill_delayer(self, delayer: PrefillDelayer | None) -> None: self.prefill_delayer = delayer + self._local_prefill_coalescing = delayer is not None and delayer.is_local + self._inflight_prefix_wait = {} + + def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: + """Bounded wait for a hybrid producer's planned prompt-end checkpoint.""" + bm = self.block_manager + delayer = self.prefill_delayer + if ( + not bm.enable_prefix_caching + or not bm.state.enabled + or bm.state_checkpoint_interval_tokens == 0 + or not seq.has_per_req_cache + or seq.multimodal_data is not None + ): + return False + deadline = self._inflight_prefix_wait.get(seq.id) + if deadline is not None and self._schedule_tick >= deadline: + return False + for producer in self.running: + anchor = producer.checkpoint_end_pos + if ( + producer.status != SequenceStatus.RUNNING + or producer.num_cached_tokens >= anchor + or anchor - cached_tokens < self.max_num_batched_tokens + or anchor >= seq.num_prompt_tokens + or producer.multimodal_data is not None + or producer.cache_seed != seq.cache_seed + ): + continue + keepers = bm.checkpointers_at( + producer, + anchor, + # A fork needs room both to preserve the producer's state and + # to restore it in the consumer's first forward. + min(producer.num_prompt_tokens - anchor, seq.num_tokens - anchor), + ) + if not keepers or any( + cache.applies(seq) and cache not in keepers for cache in bm.state_caches + ): + continue + # Reject divergent heads cheaply; equal long buffers stay in NumPy. + head = min(anchor, 64) + if ( + memoryview(seq.token_ids)[:head] + != memoryview(producer.token_ids)[:head] + ): + continue + if np.array_equal( + np.frombuffer(seq.token_ids, dtype=np.int32, count=anchor), + np.frombuffer(producer.token_ids, dtype=np.int32, count=anchor), + ): + self._inflight_prefix_wait.setdefault( + seq.id, self._schedule_tick + delayer.ttft_max_ticks + ) + return True + return False + + def _waiting_prefills(self): + for seq in self.waiting: + if ( + seq.status + not in (SequenceStatus.ABORTED, SequenceStatus.WAITING_FOR_REMOTE_KVS) + and self._unschedulable_reason(seq) is None + ): + yield seq + + def _local_prefill_pending_work(self) -> tuple[bool, int]: + """Count probed work; a bounded scan never invents fill for unseen work.""" + budget = self.max_num_batched_tokens + pending = 0 + if self._partial_prefill_count: + for seq in self.running: + if not seq.is_partial_prefill: + continue + chunk = self._partial_prefill_chunk(seq, pending) + if not chunk: + break + chunk = self._preview_prefill_chunk(seq, seq.num_cached_tokens, chunk) + pending += chunk + if pending >= budget: + break + prefillable = pending > 0 + if pending >= budget or len(self.running) >= self.max_num_seqs: + return prefillable, min(pending, budget) + probed_tokens = 0 + slots = self.max_num_seqs - len(self.running) + for seq in self._waiting_prefills(): + offload_resume = self._is_offload_prefill_resume(seq) + if offload_resume: + cached = self._offload_prefill_start(seq) + else: + if probed_tokens >= budget: + break + cached_blocks = self.block_manager.can_allocate( + seq, record=False, reuse_hashes=True + ) + probed_tokens += seq.num_tokens + if cached_blocks < 0: + break # Phase 2 also stops at the first allocation refusal. + if slots <= self._num_parked_remote_kv: + # A connector query can pin remote KV, so leave it to + # admission. Do not count uncertain remote-slot fit as + # fill or mistake an individually fitting request for idle. + prefillable = True + break + cached = cached_blocks * self.block_manager.hash_block_size + remaining = ( + seq.num_tokens - cached + if offload_resume + else self._new_prefill_tokens(seq, cached) + ) + chunk = self._prefill_chunk_for_budget(remaining, budget - pending, pending) + if chunk is None or ( + self._requires_atomic_prefill(seq) and chunk < remaining + ): + break + chunk = self._preview_prefill_chunk(seq, cached, chunk) + prefillable = True + pending += chunk + slots -= 1 + if pending >= budget: + return True, budget + if slots == 0: + break + return prefillable, pending + + def _offload_prefill_start(self, seq: Sequence) -> int: + if seq.num_cached_tokens < seq.num_tokens: + return seq.num_cached_tokens + hbs = self.block_manager.hash_block_size + return max(0, (seq.num_tokens - 1) // hbs * hbs) + + def _new_prefill_tokens(self, seq: Sequence, cached_tokens: int) -> int: + remaining = seq.num_tokens - cached_tokens + if ( + not self._requires_atomic_prefill(seq) + and self.enable_chunked_prefill + and 0 < self.long_prefill_token_threshold < remaining + ): + remaining = self.long_prefill_token_threshold + return remaining + + def _partial_prefill_chunk(self, seq: Sequence, batched_tokens: int) -> int: + remaining = seq.num_tokens - seq.num_cached_tokens + if 0 < self.long_prefill_token_threshold < remaining: + remaining = self.long_prefill_token_threshold + return self._chunked_prefill_size( + remaining, self.max_num_batched_tokens - batched_tokens, batched_tokens + ) + + def _preview_prefill_chunk(self, seq: Sequence, start: int, chunk: int) -> int: + # Estimation must not cancel a state fork. Only the actual admission + # finalizes that decision after the delayer releases the batch. + if self._requires_atomic_prefill(seq): + return chunk + target = self.block_manager.checkpoint_cut( + seq, start, start + chunk, record=False + ) + return target - start if target else chunk def _can_admit_head_prefill(self) -> bool: """Match SGL's `local_prefillable=True` semantics: report True iff @@ -732,7 +894,10 @@ def _can_admit_head_prefill(self) -> bool: break if self._unschedulable_reason(seq) is not None: continue - if seq.status == SequenceStatus.WAITING_FOR_REMOTE_KVS: + if seq.status in ( + SequenceStatus.ABORTED, + SequenceStatus.WAITING_FOR_REMOTE_KVS, + ): continue num_new_tokens = seq.num_tokens - seq.num_cached_tokens if ( @@ -803,8 +968,8 @@ def _waiting_new_token_count(self) -> int: early-exits the scan: one batch's worth is all the coalescer compares against, so there's no point summing a deep queue. - Skips the same non-admittable seqs as `_can_admit_head_prefill` — - unschedulable, WAITING_FOR_REMOTE_KVS, and oversized-when-chunking-off — + Skips unschedulable, ABORTED, WAITING_FOR_REMOTE_KVS, and + oversized-when-chunking-off sequences — so the "queued work" signal counts only tokens this rank could actually prefill this step. Counting remote-KV / unschedulable tokens here would inflate the cross-rank aggregate and reach the fill target before a real @@ -817,7 +982,10 @@ def _waiting_new_token_count(self) -> int: for seq in self.waiting: if self._unschedulable_reason(seq) is not None: continue - if seq.status == SequenceStatus.WAITING_FOR_REMOTE_KVS: + if seq.status in ( + SequenceStatus.ABORTED, + SequenceStatus.WAITING_FOR_REMOTE_KVS, + ): continue num_new_tokens = seq.num_tokens - seq.num_cached_tokens if ( @@ -868,18 +1036,19 @@ def _oldest_waiting_prefill_age_ms(self) -> float: """Age in ms (since arrival) of the oldest ADMITTABLE waiting prefill, or 0.0 if none. - Feeds PrefillDelayer's TTFT SLA guard: if this exceeds max_queue_ms the - coalescer force-releases so a request never starves in the queue. Uses - `seq.arrive_time` (wall-clock seconds, stamped at engine entry) — the - true end-to-end wait, including backlog and coalescer holds. Skips the - same non-admittable seqs as `_can_admit_head_prefill` (unschedulable, - WAITING_FOR_REMOTE_KVS) so a permanently-stuck seq can't peg the guard. + After decode protection, max_queue_ms releases extra coalescing based + on time since arrival, including backlog. This does not bound resource + or checkpoint waits. ABORTED, remote-loading and statically rejected + requests do not contribute. """ oldest_arrive = None for seq in self.waiting: if self._unschedulable_reason(seq) is not None: continue - if seq.status == SequenceStatus.WAITING_FOR_REMOTE_KVS: + if seq.status in ( + SequenceStatus.ABORTED, + SequenceStatus.WAITING_FOR_REMOTE_KVS, + ): continue if oldest_arrive is None or seq.arrive_time < oldest_arrive: oldest_arrive = seq.arrive_time @@ -1348,34 +1517,60 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: self._promote_ready_remote_kv_requests() self._park_ready_offload_partial_prefills() + # Reclaim aborted heads even when decode protection vetoes Phase 2. + while self.waiting and self.waiting[0].status == SequenceStatus.ABORTED: + self._reject_aborted_waiting(self.waiting.popleft()) # should_allow_prefill() runs a cross-DP all_reduce and MUST be called # every tick on every rank for lockstep — hence before the early-return. if self.prefill_delayer is not None: # pending = fresh waiting new-tokens + resumable partials' remaining, # capped at the batch budget: the coalescer's accumulation signal. - pending_tokens = min( - self._waiting_new_token_count() - + self._partial_prefill_remaining_tokens(), - self.max_num_batched_tokens, + running_decode_batch = max( + 0, len(self.running) - self._partial_prefill_count ) + protects_decode = self.prefill_delayer.protects_decode(running_decode_batch) + if protects_decode or ( + self._local_prefill_coalescing and not running_decode_batch + ): + # Existence is sufficient during the hard protection window: + # fit probes cannot change its decision. Coalescer hold bounds + # start afterwards. With no decode, admission proceeds directly. + prefillable = ( + self._partial_prefill_count > 0 + or next(self._waiting_prefills(), None) is not None + ) + pending_tokens = 0 + elif self._local_prefill_coalescing: + prefillable, pending_tokens = self._local_prefill_pending_work() + else: + prefillable = self._can_admit_head_prefill() + pending_tokens = min( + self._waiting_new_token_count() + + self._partial_prefill_remaining_tokens(), + self.max_num_batched_tokens, + ) delayer_allows = self.prefill_delayer.should_allow_prefill( - prefillable=self._can_admit_head_prefill(), + prefillable=prefillable, pending_tokens=pending_tokens, # decode-only: self.running also holds mid-chunked-prefill seqs, # which are NOT decode load — counting them would defeat the # coalescer's "no decode → fire" fast path. - running_decode_batch=max( - 0, len(self.running) - self._partial_prefill_count - ), - kv_usage=self._kv_usage(), + running_decode_batch=running_decode_batch, + kv_usage=0.0 if protects_decode else self._kv_usage(), has_partial=self._partial_prefill_count > 0, - oldest_waiting_age_ms=self._oldest_waiting_prefill_age_ms(), + oldest_waiting_age_ms=( + self._oldest_waiting_prefill_age_ms() + if not protects_decode + and self.prefill_delayer.max_queue_ms is not None + else 0.0 + ), ) else: delayer_allows = True - if not self.running and not self.waiting: + # Rejections may still need an empty batch to dispatch connector cleanup. + if not self.running and not self.waiting and not self._rejected: return None # ---- Phase 1: resume partial prefills from running ---- @@ -1390,13 +1585,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: break if not seq.is_partial_prefill: continue - remaining = seq.num_tokens - seq.num_cached_tokens - if 0 < self.long_prefill_token_threshold < remaining: - remaining = self.long_prefill_token_threshold - budget_remaining = self.max_num_batched_tokens - num_batched_tokens - chunk = self._chunked_prefill_size( - remaining, budget_remaining, num_batched_tokens - ) + chunk = self._partial_prefill_chunk(seq, num_batched_tokens) if chunk: chunk = self._finalize_prefill_chunk( seq, seq.num_cached_tokens, chunk @@ -1410,6 +1599,8 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: num_scheduled_tokens.append(chunk) # ---- Phase 2: new requests from waiting ---- + prefix_waiters: deque[Sequence] = deque() + prefix_bypass_left = min(16, self.max_num_seqs) while ( delayer_allows and (self.delay_factor <= 0 or self._passed_delay(time.time())) @@ -1417,6 +1608,10 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: and num_seqs_prefill < self.max_num_seqs and num_batched_tokens < self.max_num_batched_tokens ): + if prefix_waiters: + if prefix_bypass_left == 0: + break + prefix_bypass_left -= 1 seq = self.waiting.popleft() # Client disconnected before this seq ever ran: it holds no KV yet @@ -1447,6 +1642,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # Re-check here (not just at submit) since pool state may change. unschedulable = self._unschedulable_reason(seq) if unschedulable is not None: + self._inflight_prefix_wait.pop(seq.id, None) seq.status = SequenceStatus.FINISHED seq.leave_reason = f"unschedulable: {unschedulable}" seq.multimodal_data = None @@ -1504,8 +1700,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # costs one block of forward, and the forward has to happen # regardless: the first decode samples from the logits this # prefill produces, and there is nowhere else to get them. - hbs = self.block_manager.hash_block_size - seq.num_cached_tokens = max(0, (seq.num_tokens - 1) // hbs * hbs) + seq.num_cached_tokens = self._offload_prefill_start(seq) num_new_tokens = seq.num_tokens - seq.num_cached_tokens budget_remaining = self.max_num_batched_tokens - num_batched_tokens chunk = self._prefill_chunk_for_budget( @@ -1552,7 +1747,13 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # (post-prefix-cache) remaining token count. V4 SWA correctness is # enforced inside can_allocate (_swa_bounded_hit bounds the hit to # where the trailing-window SWA is present); no post-hoc warmup trim. - num_cached_blocks = self.block_manager.can_allocate(seq) + block_hashes = [] + num_cached_blocks = self.block_manager.can_allocate( + seq, + record=False, + block_hashes=block_hashes, + reuse_hashes=self._local_prefill_coalescing, + ) if num_cached_blocks < 0: self.waiting.appendleft(seq) break @@ -1561,16 +1762,23 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # their decoded tokens — preempt() frees their KV blocks but keeps # the token_ids, so num_tokens > num_prompt_tokens and those tokens # still need KV recomputed. - num_new_tokens = ( - seq.num_tokens - num_cached_blocks * self.block_manager.hash_block_size + num_new_tokens = self._new_prefill_tokens( + seq, num_cached_blocks * self.block_manager.hash_block_size ) - atomic_prefill = self._requires_atomic_prefill(seq) if ( - not atomic_prefill - and self.enable_chunked_prefill - and 0 < self.long_prefill_token_threshold < num_new_tokens + self._local_prefill_coalescing + and (not needs_remote_load or self._connector_flag("is_offload")) + and self._wait_for_inflight_prefix( + seq, + max( + num_cached_blocks * self.block_manager.hash_block_size, + seq.offload_joint.kv_prefix_tokens, + ), + ) ): - num_new_tokens = self.long_prefill_token_threshold + prefix_waiters.append(seq) + continue + atomic_prefill = self._requires_atomic_prefill(seq) budget_remaining = self.max_num_batched_tokens - num_batched_tokens chunk = self._prefill_chunk_for_budget( num_new_tokens, budget_remaining, num_batched_tokens @@ -1578,6 +1786,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: if chunk is None or (atomic_prefill and chunk < num_new_tokens): self.waiting.appendleft(seq) break + self.block_manager.record_allocation(seq, num_cached_blocks, block_hashes) if not self.block_manager.allocate(seq, num_cached_blocks): # A state-less joint boundary could not be privatised without # risking a shared decoding sequence's blocks (finding #2). @@ -1625,6 +1834,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: self.waiting.appendleft(seq) break self._park_for_remote_load(seq, skipped_waiting_requests) + self._inflight_prefix_wait.pop(seq.id, None) continue if seq.offload_joint.boundary_tokens: @@ -1661,6 +1871,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # same wake-up -- because to the scheduler both are one event: # a transfer into blocks this request already holds. self._park_for_remote_load(seq, skipped_waiting_requests) + self._inflight_prefix_wait.pop(seq.id, None) continue # Refresh, not a duplicate of the set above: that one is guarded @@ -1701,6 +1912,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: num_seqs_prefill, num_batched_tokens, ) + self._inflight_prefix_wait.pop(seq.id, None) if skipped_waiting_requests: logger.debug( @@ -1708,6 +1920,10 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: len(skipped_waiting_requests), ) self.waiting.extend(skipped_waiting_requests) + if prefix_waiters: + # Keep deferred requests in arrival order while bounded lookahead + # admits independent work without extending their deadlines. + self.waiting.extendleft(reversed(prefix_waiters)) if self._num_parked_remote_kv > 0 and self._schedule_tick % 1000 == 0: logger.info( @@ -2047,25 +2263,18 @@ def _consume_failed_remote_kv(self, seq: Sequence) -> bool: return True def _reject_aborted_waiting(self, seq: Sequence) -> None: + self._inflight_prefix_wait.pop(seq.id, None) has_inflight_load = bool(getattr(seq, "_counted_as_inflight_load", False)) seq.status = SequenceStatus.FINISHED seq.leave_reason = "aborted" seq.multimodal_data = None self._rejected.append(seq) - if not has_inflight_load and self.kv_connector is not None: - # This path bypasses postprocess: producers must discard queued - # saves, and offload may own CPU pins before HBM allocation. - # Already-dispatched loads retain the completion-driven path below. + if not has_inflight_load: + # A lookup can pin CPU KV before HBM allocation succeeds. if self._connector_flag("is_offload"): self.kv_connector.cancel_pending_load(seq) - self.kv_connector.request_finished(seq) - # No send will ever be issued for an abort, so nothing would - # otherwise retire a claim taken above. See `postprocess`. - self._connector_send_finished(seq.id) self.deferred_free_blocks[seq.id] = seq - self._maybe_release_deferred(seq) - if not has_inflight_load or not self._connector_flag("is_offload"): - self._uncount_inflight_load(seq) + self._cleanup_aborted_load(seq) return self.deferred_free_blocks[seq.id] = seq @@ -2081,6 +2290,20 @@ def _reject_aborted_waiting(self, seq: Sequence) -> None: return seq._awaiting_aborted_load_cleanup = True + def abort_request(self, req_id: int) -> bool: + """Reclaim waiting cancellations on the event, independent of admission.""" + for seq in self.waiting: + if seq.id == req_id: + self.waiting.remove(seq) + seq.status = SequenceStatus.ABORTED + self._reject_aborted_waiting(seq) + return True + for seq in self.running: + if seq.id == req_id: + seq.status = SequenceStatus.ABORTED + return True + return False + def _cleanup_aborted_load(self, seq: Sequence) -> None: if hasattr(seq, "_awaiting_aborted_load_cleanup"): delattr(seq, "_awaiting_aborted_load_cleanup") @@ -3512,10 +3735,12 @@ def _update_from_kv_xfer_finished(self, kv_connector_output: KVConnectorOutput): kv_connector_output = process_completions(kv_connector_output) for req_id in kv_connector_output.finished_recving or (): - if metrics := getattr(self, "metrics", None): - metrics.finish_kv_wait(req_id, succeeded=True) assert not is_producer, "Only consumer should update recving KV status" logger.debug("Finished recving KV transfer for request %s", req_id) + if self._finish_aborted_load_cleanup(req_id): + continue + if metrics := getattr(self, "metrics", None): + metrics.finish_kv_wait(req_id, succeeded=True) self.finished_recving_kv_req_ids.append(req_id) for req_id in kv_connector_output.failed_recving or (): @@ -3525,6 +3750,8 @@ def _update_from_kv_xfer_finished(self, kv_connector_output: KVConnectorOutput): logger.warning( "KV receive failed for request %s; falling back to prefill.", req_id ) + if self._finish_aborted_load_cleanup(req_id): + continue self.failed_recving_kv_req_ids.append(req_id) # The two loading channels carry state-tier loads as well as KV ones, @@ -3672,7 +3899,8 @@ def get_next_batch_info(self) -> tuple[bool, int, int]: eligible_waiting = [ seq for seq in self.waiting - if seq.status != SequenceStatus.WAITING_FOR_REMOTE_KVS + if seq.status + not in (SequenceStatus.ABORTED, SequenceStatus.WAITING_FOR_REMOTE_KVS) ] if eligible_waiting: # new request is waiting, will do prefill @@ -3945,6 +4173,16 @@ class DecodeScheduler(Scheduler): _ENGINE_LABEL = "Decode " _METRICS_ROLE = "decode" + def abort_request(self, req_id: int) -> bool: + # RapidServe does not drain the standard scheduler's rejection queue. + # Preserve its mark-only cancellation path. + with self._prefill_lock: + for seq in (*self.running, *self.waiting): + if seq.id == req_id: + seq.status = SequenceStatus.ABORTED + return True + return False + def get_request_counts(self) -> tuple[int, int]: """Fold in the two queues this scheduler adds. diff --git a/atom/utils/envs.py b/atom/utils/envs.py index 4b65fdc2e3..30b518e7ff 100644 --- a/atom/utils/envs.py +++ b/atom/utils/envs.py @@ -721,11 +721,9 @@ def _positive_float_env(name: str, default: str) -> float: if os.getenv("ATOM_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK", "") == "" else float(os.getenv("ATOM_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK")) ), - # TTFT SLA guard: if any rank's oldest schedulable waiting prefill has queued - # (since arrival) >= this many ms, force-release regardless of the fill - # target. Bounds worst-case TTFT. Empty string => None => disabled (set this - # to your TTFT budget in ms to activate; a small value under heavy backlog - # will fire every tick and defeat coalescing, so size it to the SLA). + # After decode protection, bound extra coalescing by queue age. Checkpoint + # dependency waits use TTFT_MAX_TICKS; this is not an end-to-end TTFT bound. + # Empty string => None => disabled. "ATOM_PREFILL_DELAYER_MAX_QUEUE_MS": lambda: ( None if os.getenv("ATOM_PREFILL_DELAYER_MAX_QUEUE_MS", "") == "" @@ -734,6 +732,7 @@ def _positive_float_env(name: str, default: str) -> float: # After a prefill forward, protect this many scheduler passes for decode # before allowing another prefill. Mirrors SGLang's # --prefill-decode-interval; 0 disables the hard interval. + # A nonzero interval also enables local coalescing on TP without PP. "ATOM_PREFILL_DECODE_INTERVAL": lambda: int( os.getenv("ATOM_PREFILL_DECODE_INTERVAL", "0") ), diff --git a/docs/environment_variables.md b/docs/environment_variables.md index 9a3c11eec4..ae7e442d61 100644 --- a/docs/environment_variables.md +++ b/docs/environment_variables.md @@ -14,19 +14,46 @@ This document describes the environment variables used in the ATOM project. | **ATOM_DP_LB_REQ_EQUIV** | int | 512 | Token-equivalent decode pressure assigned to each in-flight request by `least_tokens` routing. | | **ATOM_DP_SESSION_AFFINITY** | bool | false | Load-place each new session, then keep later turns on the same prefix-cache owner. Reads `X-Dynamo-Session-ID`, falling back to `X-Correlation-ID`. | -## Prefill delayer (DP attention) - -Prefill **coalescer** for DP-attention + EP-MoE serving. Holds back prefill -admission until the accumulated prefill (fresh waiting tokens + resumable -partials' remaining tokens) fills a worthwhile forward, so fragmented -short-input prefills / small partial tail chunks batch into one forward instead -of firing many tiny ones. Releases when the fill target is reached, when a -must-fire bound trips (no decode to hide behind, KV pressure/starvation, TTFT -deadline, partial deadline), or when the queue stops growing. Preserves -cross-rank phase alignment (releases only when every rank is prefill-ready, -unless a bound forces it). All timing is tick-based (deterministic across ranks — -no wall-clock skew). See `atom/model_engine/prefill_delayer.py`. Active only when -`data_parallel_size > 1`. +## Prefill delayer (TP/DCP and DP attention) + +Coalesces waiting prefills while decode continues. DP attention enables it by +default through `ATOM_ENABLE_PREFILL_DELAYER`. For a single scheduler (DP=1, +PP=1), including TP/DCP, it is opt-in: set `ATOM_PREFILL_DECODE_INTERVAL` above +zero and keep the master switch enabled. Interval 0 leaves TP scheduling +unchanged; setting only the master switch does not enable TP coalescing. +On TP, the interval and coalescer are enabled together. +This applies to the standard scheduler, including connector-based P/D roles. +RapidServe's dedicated `PrefillScheduler`/`DecodeScheduler` do not use the delayer. + +After each executed prefill, the decode interval runs before all coalescing +bounds. Once it expires, fill, queue age, KV pressure, partial-prefill and stall +bounds decide when to release. `MAX_QUEUE_MS` stops extra coalescing after that +interval; it does not guarantee end-to-end TTFT. DP decisions reduce local +signals across ranks to keep their phases aligned. + +The local fill signal discounts HBM cache hits and uses the admission path's +chunk limits. To bound CPU work, it stops probing fresh requests after their +total prompt length reaches one batch budget. It reports only work found so +far; it does not treat unseen requests as a full batch. A deep, cache-heavy +queue can therefore release through the stall or hold bounds before reaching +the fill target. If parked transfers exhaust the unreserved slots, a fresh +request may signal possible work with zero estimated tokens until admission +resolves its connector match. This is an estimate, not a reservation: checkpoint dependencies, +connector results and resource changes during admission can still reduce a batch. +Repeated probes reuse immutable prompt hashes while rechecking pool contents +and resource fit. During decode protection, only the existence of queued or +partial work is checked; partial-prefill hold bounds start after the interval. + +Local hybrid models with state checkpointing can wait for an in-flight +producer's reusable prompt-end checkpoint. The expected prefix must exceed both +the HBM hit and any offload match by at least one prefill budget. P/D transfers +and already-started offload loads keep their own progress paths. Deferred +requests retain their relative order, while at most 16 later queue entries are +examined for independent work each pass. Each request's wait expires after +`TTFT_MAX_TICKS` scheduler passes from its first dependency wait; bypassing it +does not restart that deadline. `MAX_QUEUE_MS` bounds coalescing, not this +dependency wait: time spent queued before a producer becomes runnable does not +make duplicate prefill useful. Pure-attention models do not use checkpoint waits. | Variable | Type | Default | Description | |----------|------|---------|-------------| @@ -37,8 +64,8 @@ no wall-clock skew). See `atom/model_engine/prefill_delayer.py`. Active only whe | **ATOM_PREFILL_DELAYER_STALL_TICKS** | int | 10 | After this many consecutive non-growing ticks, release (burst ended, more won't come). Values `< 1` clamped to 1. | | **ATOM_PREFILL_DELAYER_KV_HIGH_WATERMARK** | float | 0.9 | At/above this KV usage a prefillable rank force-releases (can't accumulate a bigger batch anyway). | | **ATOM_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK** | float\|"" | "" (None) | If set, a prefillable rank below this KV usage force-releases (GPU starving). | -| **ATOM_PREFILL_DELAYER_MAX_QUEUE_MS** | float\|"" | "" (None) | TTFT SLA guard: if any rank's oldest schedulable waiting prefill has queued (since arrival) ≥ this many ms, force-release regardless of the fill target. Measures true end-to-end wait (backlog + coalescer holds), unlike the tick-based TTFT bound which only caps one hold episode. Empty = disabled; set to your TTFT budget (a small value under heavy backlog fires every tick and defeats coalescing). | -| **ATOM_PREFILL_DECODE_INTERVAL** | int | 0 | After an executed prefill forward, protect this many scheduler passes for decode before admitting another prefill. `0` disables the interval. | +| **ATOM_PREFILL_DELAYER_MAX_QUEUE_MS** | float\|"" | "" (None) | After decode protection, release coalescing when the oldest schedulable waiting prefill reaches this age since arrival. Empty disables the age guard. Checkpoint dependency waits use `TTFT_MAX_TICKS`. This is not a hard TTFT limit. | +| **ATOM_PREFILL_DECODE_INTERVAL** | int | 0 | Protect this many scheduler passes after an executed prefill. On DP=1, PP=1, a positive value also enables local coalescing when the master switch is on; `0` leaves TP scheduling unchanged. On DP>1, `0` disables only the interval. | | **ATOM_PREFILL_DELAYER_DEBUG** | bool | false | Per-tick FIRE/HOLD debug logging. | | **ATOM_PREFILL_DELAYER_LOG_EVERY** | int | 1000 | Emit aggregate stats (per-exit fire counts + hold rate) every N decisions (0 disables). | diff --git a/docs/scheduling_kv_cache_guide.md b/docs/scheduling_kv_cache_guide.md index cf9ac10c91..8b9fbf12a7 100644 --- a/docs/scheduling_kv_cache_guide.md +++ b/docs/scheduling_kv_cache_guide.md @@ -62,25 +62,29 @@ The scheduler maintains two deques — `waiting` (pending prefill) and `running` ### Schedule flow -`Scheduler.schedule()` proceeds in two phases: - -**Phase 1 — Prefill scheduling:** - -1. While the delay gate passes (`_passed_delay`), the waiting queue is non-empty, and `num_seqs_prefill < max_num_seqs`: - - Peek the first waiting sequence. - - Compute `num_new_tokens = seq.num_tokens - seq.num_cached_tokens` (prefix cache hits reduce new tokens). - - If `num_batched_tokens + num_new_tokens > max_num_batched_tokens` or `block_manager.can_allocate(seq)` returns `False`, break. - - Otherwise: allocate blocks, set `seq.status = RUNNING`, `seq.type = PREFILL`, move from `waiting` to `running`. -2. If any prefill sequences were scheduled, return the batch immediately (no decode mixing). - -**Phase 2 — Decode scheduling (only when zero prefills were scheduled):** - -1. Pop sequences from `running` up to `max_num_seqs`. -2. For each sequence, check `block_manager.can_append(seq)`. -3. If a block cannot be appended, **preempt** the last running sequence (move it back to `waiting` with status `WAITING` and deallocate its blocks). -4. If the sequence has speculative draft tokens (`seq.spec_token_ids`), record them in `scheduled_spec_decode_tokens`. -5. Call `block_manager.may_append(seq, num_new_tokens)` where `num_new_tokens = mtp_k + 1`. -6. Re-insert all scheduled sequences back into `running` (preserving order). +`Scheduler.schedule()` first promotes completed transfers and asks the prefill +delayer whether this tick allows prefill. It then selects work in three phases: + +1. Resume partial prefills from `running`, subject to the token budget, + `long_prefill_token_threshold` and state-checkpoint chunk boundaries. +2. Admit fresh or offload-resumed prefills from `waiting`. A fresh request calls + `can_allocate(seq, record=False, block_hashes=hashes)` to obtain its HBM hit + and check capacity (`-1` means refusal; `0` is a valid cold admission). + After dependency-wait and token-budget checks, `record_allocation` commits + the joint boundary and `allocate` claims resources. Remote loads park until + their completion; ready offload resumes retain their existing allocations. + Local checkpoint waiters permit a bounded scan for independent work and + retain their relative FIFO order. See the [delayer settings](environment_variables.md#prefill-delayer-tpdcp-and-dp-attention). +3. If no prefill was selected, schedule decode from `running`. `can_append` + checks extension capacity; pressure can preempt a running request back to + `waiting`. `may_append` reserves the selected tokens, including speculative + tokens when enabled. Partial prefills skipped by the delayer return to the + tail of `running`. + +Prefill batches return without decode mixing. A waiting cancellation is removed +on the abort event rather than waiting for admission to reach it. In-flight +transfers retain their resources until their terminal notification; running +cancellations finish through postprocess. ### Delay factor @@ -404,30 +408,29 @@ A checkpoint the seq resumed from is not released here — it went back to the f ### Can-allocate and can-append checks ```python -def can_allocate(self, seq: Sequence) -> int: - """Return the number of cache-hit blocks (>=0) if seq fits, else -1.""" - # State cache has its own reservation; admission only needs a free slot - # index, not extra paged blocks. - if seq.has_per_req_cache and not self.state.has_free(): - return -1 - if not self.enable_prefix_caching: - if not self.kv.has_free(self.num_pool_blocks(len(seq))): - return -1 - # ... (prefix caching dry-run returns the contiguous hit-block count) +def can_allocate(self, seq: Sequence, record: bool = True, *, + block_hashes: list[int] | None = None, + reuse_hashes: bool = False) -> int: + """Return cache-hit hash blocks (>=0) if seq fits, else -1.""" + +def record_allocation(self, seq: Sequence, num_cached_blocks: int, + block_hashes: list[int]) -> None: + """Commit the joint boundary after the scheduler accepts the fit probe.""" def can_append(self, seq: Sequence, num_new_tokens: int = 1) -> bool: - seq_len = len(seq) - current_blocks = len(seq.block_table) - needed_blocks = (seq_len + num_new_tokens + self.block_size - 1) // self.block_size - new_blocks_needed = max(0, needed_blocks - current_blocks) - return self.kv.has_free(new_blocks_needed) + """Check capacity for the selected decode extension.""" ``` -- `can_allocate` checks that: - - Enough free KV blocks exist for the full sequence. A windowed architecture adds nothing here: its window is a ring inside the per-request state slot, so the slot check below covers it. - - At least one per-request cache slot group is available if the sequence has `has_per_req_cache=True`. Per-request state costs no paged blocks — its bytes were reserved ahead of the paged pool at sizing time. - -- `can_append` checks whether a decode step needs a new block. Calculates the required block count given `num_new_tokens` (typically `mtp_k + 1` for speculative decode) and returns whether enough free blocks remain. +`can_allocate` checks the full per-request state-slot reservation and PAGE-unit +capacity. Its cache hit incorporates KV, SWA and state-checkpoint availability. +With `record=False`, it leaves joint-load fields and joint-boundary counters +uncommitted; checkpoint demand/end and hit instrumentation still refresh. +Checkpoint demand counters deduplicate per request. With `reuse_hashes=True`, +immutable prompt hashes are memoized, but pool lookup, token equality and fit +are checked on every call. Mutable completion tokens are never memoized. + +`can_append` checks whether the next decode extension needs additional blocks, +including speculative tokens and the active cache layout. ### May-append (decode extension) diff --git a/tests/test_dense_offload_connector.py b/tests/test_dense_offload_connector.py index 30fb67f26e..792743f94a 100644 --- a/tests/test_dense_offload_connector.py +++ b/tests/test_dense_offload_connector.py @@ -507,7 +507,7 @@ def lookup(_tokens, lookup_id): scheduler.add(seq) allocation_attempts = [] - def cannot_allocate(value): + def cannot_allocate(value, **_kwargs): allocation_attempts.append(value.id) return -1 diff --git a/tests/test_lmcache_offload_connector.py b/tests/test_lmcache_offload_connector.py index dade09ca31..0477dff6a2 100644 --- a/tests/test_lmcache_offload_connector.py +++ b/tests/test_lmcache_offload_connector.py @@ -3584,6 +3584,7 @@ def deallocate(value): value.state_slot = -1 host = Scheduler.__new__(Scheduler) + host._inflight_prefix_wait = {} host.kv_connector = _Connector() host.block_manager = SimpleNamespace(deallocate=deallocate) host.deferred_free_blocks = {} @@ -3650,6 +3651,7 @@ def deallocate(value): value.state_slot = -1 host = Scheduler.__new__(Scheduler) + host._inflight_prefix_wait = {} host.kv_connector = _Connector() host.block_manager = SimpleNamespace(deallocate=deallocate) host.deferred_free_blocks = {} @@ -3723,6 +3725,7 @@ def deallocate(value): value.state_slot = -1 host = Scheduler.__new__(Scheduler) + host._inflight_prefix_wait = {} host.kv_connector = _Connector() host.block_manager = SimpleNamespace( kv_events_enabled=False, diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index cb8b4a1b03..0ae1bb0642 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -428,7 +428,7 @@ def test_get_request_counts(self, scheduler, seq_factory): class TestSchedule: - def test_non_offload_abort_keeps_existing_receive_cleanup(self): + def test_non_offload_abort_retains_receive_until_terminal(self): seq = SimpleNamespace( id=96, status=SequenceStatus.ABORTED, @@ -436,6 +436,7 @@ def test_non_offload_abort_keeps_existing_receive_cleanup(self): ) sched = Scheduler.__new__(Scheduler) sched._rejected = [] + sched._inflight_prefix_wait = {} sched.deferred_free_blocks = {} sched.finished_recving_kv_req_ids = [] sched.failed_recving_kv_req_ids = [] @@ -444,8 +445,9 @@ def test_non_offload_abort_keeps_existing_receive_cleanup(self): sched._reject_aborted_waiting(seq) - assert sched.deferred_free_blocks == {} - assert sched._num_parked_remote_kv == 0 + assert sched.deferred_free_blocks == {seq.id: seq} + assert sched._num_parked_remote_kv == 1 + assert seq._awaiting_aborted_load_cleanup assert sched._rejected == [seq] @pytest.mark.parametrize("producer", [False, True]) @@ -508,6 +510,7 @@ def source_blocks_released(self, value): sched.waiting = deque() sched.running = deque() sched._rejected = [] + sched._inflight_prefix_wait = {} sched.deferred_free_blocks = {} sched.finished_recving_kv_req_ids = [] sched.failed_recving_kv_req_ids = [] diff --git a/tests/test_scheduler_partial_prefill_tail.py b/tests/test_scheduler_partial_prefill_tail.py index eefebaed30..3f3183eeff 100644 --- a/tests/test_scheduler_partial_prefill_tail.py +++ b/tests/test_scheduler_partial_prefill_tail.py @@ -45,6 +45,12 @@ class _VetoDelayer: """Stub cross-DP delayer that always refuses prefill this tick, forcing the decode loop to run while a partial prefill is still sitting in `running`.""" + is_local = False + max_queue_ms = None + + def protects_decode(self, running_decode_batch: int) -> bool: + return False + def should_allow_prefill(self, prefillable, pending_tokens, **kwargs): return False