From a6864d4992e52b6febaa4c7aa365e6c7265413d0 Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 15 Sep 2026 09:06:38 +0000 Subject: [PATCH 1/8] perf(scheduler): coalesce agentic prefills on TP --- atom/model_engine/engine_core.py | 46 +++++++------ atom/model_engine/prefill_delayer.py | 15 +++++ atom/model_engine/scheduler.py | 97 +++++++++++++++++++++++++--- atom/utils/envs.py | 1 + 4 files changed, 131 insertions(+), 28 deletions(-) diff --git a/atom/model_engine/engine_core.py b/atom/model_engine/engine_core.py index 320415a763..b3851b7b6d 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,27 @@ 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: + 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 +753,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/prefill_delayer.py b/atom/model_engine/prefill_delayer.py index fdebd50255..f67be536bc 100644 --- a/atom/model_engine/prefill_delayer.py +++ b/atom/model_engine/prefill_delayer.py @@ -250,6 +250,21 @@ def _clamp_ticks(name: str, value: int) -> int: return 1 return value + def protects_decode(self, running_decode_batch: int) -> bool: + """Whether the local decision can skip cache probes during protection.""" + return ( + self.cpu_group is None + 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, diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index 9bc823a3bb..e5e9c506c4 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -630,6 +630,9 @@ def __init__( kv_events_cfg = getattr(config, "kv_events_config", None) parallel_cfg = getattr(config, "parallel_config", None) + self._local_prefill_coalescing = ( + getattr(parallel_cfg, "data_parallel_size", 1) == 1 + ) dp_rank = ( getattr(parallel_cfg, "data_parallel_rank", None) if parallel_cfg is not None @@ -698,8 +701,7 @@ def __init__( endpoint="", ) - # Cross-DP prefill alignment. Set by DPEngineCoreProc after - # dp_group is available. See `prefill_delayer.py` for rationale. + # Set by EngineCore for cross-DP alignment or opt-in TP decode protection. from atom.model_engine.prefill_delayer import PrefillDelayer self.prefill_delayer: PrefillDelayer | None = None @@ -707,6 +709,62 @@ def __init__( def set_prefill_delayer(self, delayer) -> None: self.prefill_delayer = delayer + def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: + """Avoid recomputing a prefix an admitted prefill will checkpoint.""" + if ( + not self.block_manager.enable_prefix_caching + or seq.multimodal_data is not None + ): + return False + for producer in self.running: + anchor = producer.checkpoint_end_pos + if ( + producer.status != SequenceStatus.RUNNING + or producer.num_cached_tokens >= producer.num_prompt_tokens + or anchor - cached_tokens < self.max_num_batched_tokens + or anchor >= seq.num_prompt_tokens + or producer.multimodal_data is not None + ): + continue + if seq.token_ids[:anchor] == producer.token_ids[:anchor]: + return True + return False + + def _local_prefill_pending_work(self) -> tuple[bool, int]: + """Estimate work after HBM reuse, with bounded queue probes.""" + budget = self.max_num_batched_tokens + pending = self._partial_prefill_remaining_tokens() + prefillable = pending > 0 + if pending >= budget or len(self.running) >= self.max_num_seqs: + return prefillable, min(pending, budget) + for i, seq in enumerate(self.waiting): + if i >= 4: + break + if ( + seq.status + in ( + SequenceStatus.ABORTED, + SequenceStatus.WAITING_FOR_REMOTE_KVS, + ) + or self._unschedulable_reason(seq) is not None + ): + continue + if self._is_offload_prefill_resume(seq): + cached = seq.num_cached_tokens + else: + cached_blocks = self.block_manager.can_allocate(seq, record=False) + if cached_blocks < 0: + continue + cached = cached_blocks * self.block_manager.hash_block_size + remaining = seq.num_tokens - cached + if not self.enable_chunked_prefill and remaining > budget: + continue + prefillable = True + pending += max(0, remaining) + if pending >= budget: + return True, budget + return prefillable, pending + def _can_admit_head_prefill(self) -> bool: """Match SGL's `local_prefillable=True` semantics: report True iff this rank would *actually* admit a new prefill this tick. @@ -1354,20 +1412,29 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: 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 = getattr(self.prefill_delayer, "protects_decode", None) + if protects_decode is not None and protects_decode(running_decode_batch): + prefillable = bool(self.waiting) or self._partial_prefill_count > 0 + 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 - ), + running_decode_batch=running_decode_batch, kv_usage=self._kv_usage(), has_partial=self._partial_prefill_count > 0, oldest_waiting_age_ms=self._oldest_waiting_prefill_age_ms(), @@ -1564,6 +1631,16 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: num_new_tokens = ( seq.num_tokens - num_cached_blocks * self.block_manager.hash_block_size ) + if ( + self._local_prefill_coalescing + and self.prefill_delayer is not None + and self._wait_for_inflight_prefix( + seq, num_cached_blocks * self.block_manager.hash_block_size + ) + ): + # Wait without holding blocks; re-probe after producer progress. + skipped_waiting_requests.append(seq) + continue atomic_prefill = self._requires_atomic_prefill(seq) if ( not atomic_prefill diff --git a/atom/utils/envs.py b/atom/utils/envs.py index 4b65fdc2e3..834e08147a 100644 --- a/atom/utils/envs.py +++ b/atom/utils/envs.py @@ -734,6 +734,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") ), From 0c13ac7b8307d331ed7546bebb9dcbc507eede7f Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 15 Sep 2026 09:43:56 +0000 Subject: [PATCH 2/8] refactor(scheduler): call decode protection directly --- atom/model_engine/scheduler.py | 3 +-- tests/test_scheduler_partial_prefill_tail.py | 3 +++ 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index e5e9c506c4..3d7472cd1d 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -1415,8 +1415,7 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: running_decode_batch = max( 0, len(self.running) - self._partial_prefill_count ) - protects_decode = getattr(self.prefill_delayer, "protects_decode", None) - if protects_decode is not None and protects_decode(running_decode_batch): + if self.prefill_delayer.protects_decode(running_decode_batch): prefillable = bool(self.waiting) or self._partial_prefill_count > 0 pending_tokens = 0 elif self._local_prefill_coalescing: diff --git a/tests/test_scheduler_partial_prefill_tail.py b/tests/test_scheduler_partial_prefill_tail.py index eefebaed30..79ee337929 100644 --- a/tests/test_scheduler_partial_prefill_tail.py +++ b/tests/test_scheduler_partial_prefill_tail.py @@ -45,6 +45,9 @@ 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`.""" + def protects_decode(self, running_decode_batch: int) -> bool: + return False + def should_allow_prefill(self, prefillable, pending_tokens, **kwargs): return False From a435855021041c5bce575b547137d385e308576e Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 15 Sep 2026 10:14:11 +0000 Subject: [PATCH 3/8] refactor(scheduler): track active local prefill coalescing Update the local coalescing flag when installing or removing a delayer, and compare shared token prefixes through NumPy buffer views to avoid copying both arrays. Keep the existing cross-DP test stub explicit about its topology. --- atom/model_engine/scheduler.py | 20 ++++++++++---------- tests/test_scheduler_partial_prefill_tail.py | 2 ++ 2 files changed, 12 insertions(+), 10 deletions(-) diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index 3d7472cd1d..ee39d4eaf8 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -630,9 +630,6 @@ def __init__( kv_events_cfg = getattr(config, "kv_events_config", None) parallel_cfg = getattr(config, "parallel_config", None) - self._local_prefill_coalescing = ( - getattr(parallel_cfg, "data_parallel_size", 1) == 1 - ) dp_rank = ( getattr(parallel_cfg, "data_parallel_rank", None) if parallel_cfg is not None @@ -705,9 +702,13 @@ def __init__( from atom.model_engine.prefill_delayer import PrefillDelayer self.prefill_delayer: PrefillDelayer | None = None + self._local_prefill_coalescing = False def set_prefill_delayer(self, delayer) -> None: self.prefill_delayer = delayer + self._local_prefill_coalescing = ( + delayer is not None and delayer.dp_size == 1 and delayer.cpu_group is None + ) def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: """Avoid recomputing a prefix an admitted prefill will checkpoint.""" @@ -726,7 +727,10 @@ def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: or producer.multimodal_data is not None ): continue - if seq.token_ids[:anchor] == producer.token_ids[:anchor]: + if np.array_equal( + np.frombuffer(seq.token_ids, dtype=np.int32, count=anchor), + np.frombuffer(producer.token_ids, dtype=np.int32, count=anchor), + ): return True return False @@ -1630,12 +1634,8 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: num_new_tokens = ( seq.num_tokens - num_cached_blocks * self.block_manager.hash_block_size ) - if ( - self._local_prefill_coalescing - and self.prefill_delayer is not None - and self._wait_for_inflight_prefix( - seq, num_cached_blocks * self.block_manager.hash_block_size - ) + if self._local_prefill_coalescing and self._wait_for_inflight_prefix( + seq, num_cached_blocks * self.block_manager.hash_block_size ): # Wait without holding blocks; re-probe after producer progress. skipped_waiting_requests.append(seq) diff --git a/tests/test_scheduler_partial_prefill_tail.py b/tests/test_scheduler_partial_prefill_tail.py index 79ee337929..aa744953d9 100644 --- a/tests/test_scheduler_partial_prefill_tail.py +++ b/tests/test_scheduler_partial_prefill_tail.py @@ -45,6 +45,8 @@ 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`.""" + dp_size = 2 + def protects_decode(self, running_decode_batch: int) -> bool: return False From 07b9fdc518edeb95aeb3d3f3b44c1257a6f1ed12 Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Mon, 21 Sep 2026 13:44:20 +0000 Subject: [PATCH 4/8] fix(scheduler): bound local prefill waits and probes --- atom/model_engine/block_manager.py | 37 +++-- atom/model_engine/prefill_delayer.py | 36 +++-- atom/model_engine/scheduler.py | 158 ++++++++++++++----- docs/environment_variables.md | 37 +++-- tests/test_dense_offload_connector.py | 2 +- tests/test_scheduler_partial_prefill_tail.py | 3 +- 6 files changed, 195 insertions(+), 78 deletions(-) diff --git a/atom/model_engine/block_manager.py b/atom/model_engine/block_manager.py index 3e96ab7ccf..2cb8603655 100644 --- a/atom/model_engine/block_manager.py +++ b/atom/model_engine/block_manager.py @@ -861,17 +861,20 @@ 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, + ) -> 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 funnel counters. It still refreshes checkpoint + demand/end and hit instrumentation. 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 +891,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,7 +906,6 @@ 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] = [] 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) @@ -966,8 +972,7 @@ def can_allocate(self, seq: Sequence, record: bool = True) -> int: # 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. 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 +988,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). diff --git a/atom/model_engine/prefill_delayer.py b/atom/model_engine/prefill_delayer.py index f67be536bc..546923bb75 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,10 +258,14 @@ 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.cpu_group is None + self.is_local and not self._first and running_decode_batch > 0 and ( @@ -297,8 +309,8 @@ def should_allow_prefill( release once held for partial_max_ticks (it holds KV). 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). @@ -368,8 +380,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: @@ -392,10 +404,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 ee39d4eaf8..abc8532e43 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -703,70 +703,124 @@ def __init__( self.prefill_delayer: PrefillDelayer | None = None self._local_prefill_coalescing = False + self._inflight_prefix_wait: tuple[int, int] | None = None def set_prefill_delayer(self, delayer) -> None: self.prefill_delayer = delayer - self._local_prefill_coalescing = ( - delayer is not None and delayer.dp_size == 1 and delayer.cpu_group is None - ) + self._local_prefill_coalescing = delayer is not None and delayer.is_local + self._inflight_prefix_wait = None def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: - """Avoid recomputing a prefix an admitted prefill will checkpoint.""" + """Bounded wait for a hybrid producer's planned prompt-end checkpoint.""" + bm = self.block_manager + delayer = self.prefill_delayer if ( - not self.block_manager.enable_prefix_caching + 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 + if self._inflight_prefix_wait is not None: + waiting_id, deadline = self._inflight_prefix_wait + if waiting_id == seq.id and self._schedule_tick >= deadline: + return False + if ( + delayer.max_queue_ms is not None + and (time.time() - seq.arrive_time) * 1000 >= delayer.max_queue_ms + ): + return False for producer in self.running: anchor = producer.checkpoint_end_pos if ( producer.status != SequenceStatus.RUNNING - or producer.num_cached_tokens >= producer.num_prompt_tokens + 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, + 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), ): + if ( + self._inflight_prefix_wait is None + or self._inflight_prefix_wait[0] != seq.id + ): + self._inflight_prefix_wait = ( + 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]: - """Estimate work after HBM reuse, with bounded queue probes.""" + """Estimate HBM-discounted work; cap further probes after one prompt budget.""" budget = self.max_num_batched_tokens pending = self._partial_prefill_remaining_tokens() prefillable = pending > 0 if pending >= budget or len(self.running) >= self.max_num_seqs: return prefillable, min(pending, budget) - for i, seq in enumerate(self.waiting): - if i >= 4: - break - if ( - seq.status - in ( - SequenceStatus.ABORTED, - SequenceStatus.WAITING_FOR_REMOTE_KVS, - ) - or self._unschedulable_reason(seq) is not None - ): - continue + probed_tokens = 0 + slots = self.max_num_seqs - len(self.running) + for seq in self._waiting_prefills(): if self._is_offload_prefill_resume(seq): - cached = seq.num_cached_tokens + # Already owns its KV and state slots. Mirror the one-block + # recompute used by Phase 2 when the tier covers the whole seq. + cached = min( + seq.num_cached_tokens, + (seq.num_tokens - 1) + // self.block_manager.hash_block_size + * self.block_manager.hash_block_size, + ) else: + if probed_tokens >= budget: + # A deep cache-heavy queue must not turn into a fixed stall + # delay just because estimating its fill is expensive. + return prefillable, budget if prefillable else pending cached_blocks = self.block_manager.can_allocate(seq, record=False) + probed_tokens += seq.num_tokens if cached_blocks < 0: - continue + break # Phase 2 also stops at the first allocation refusal. cached = cached_blocks * self.block_manager.hash_block_size remaining = seq.num_tokens - cached if not self.enable_chunked_prefill and remaining > budget: continue prefillable = True - pending += max(0, remaining) + pending += remaining + slots -= 1 if pending >= budget: return True, budget + if slots == 0: + break return prefillable, pending def _can_admit_head_prefill(self) -> bool: @@ -794,7 +848,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 ( @@ -879,7 +936,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 ( @@ -941,7 +1001,10 @@ def _oldest_waiting_prefill_age_ms(self) -> float: 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 @@ -1410,6 +1473,9 @@ 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. @@ -1419,8 +1485,14 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: running_decode_batch = max( 0, len(self.running) - self._partial_prefill_count ) - if self.prefill_delayer.protects_decode(running_decode_batch): - prefillable = bool(self.waiting) or self._partial_prefill_count > 0 + protects_decode = self.prefill_delayer.protects_decode(running_decode_batch) + if protects_decode or ( + self._local_prefill_coalescing and not running_decode_batch + ): + 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() @@ -1438,14 +1510,20 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # which are NOT decode load — counting them would defeat the # coalescer's "no decode → fire" fast path. running_decode_batch=running_decode_batch, - kv_usage=self._kv_usage(), + 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 ---- @@ -1622,7 +1700,10 @@ 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 + ) if num_cached_blocks < 0: self.waiting.appendleft(seq) break @@ -1634,12 +1715,16 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: num_new_tokens = ( seq.num_tokens - num_cached_blocks * self.block_manager.hash_block_size ) - if self._local_prefill_coalescing and self._wait_for_inflight_prefix( - seq, num_cached_blocks * self.block_manager.hash_block_size + if ( + self._local_prefill_coalescing + and not needs_remote_load + and not seq.offload_joint.kv_prefix_tokens + and self._wait_for_inflight_prefix( + seq, num_cached_blocks * self.block_manager.hash_block_size + ) ): - # Wait without holding blocks; re-probe after producer progress. - skipped_waiting_requests.append(seq) - continue + self.waiting.appendleft(seq) + break atomic_prefill = self._requires_atomic_prefill(seq) if ( not atomic_prefill @@ -1654,6 +1739,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). diff --git a/docs/environment_variables.md b/docs/environment_variables.md index 9a3c11eec4..115bc29e1a 100644 --- a/docs/environment_variables.md +++ b/docs/environment_variables.md @@ -14,19 +14,26 @@ 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. + +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. + +Local hybrid models with state checkpointing also wait briefly for an in-flight +producer's reusable prompt-end checkpoint. This preserves queue order, bypasses +requests with remote KV matches, and expires after `TTFT_MAX_TICKS` scheduler +passes or `MAX_QUEUE_MS` since arrival, whichever applies first. Pure-attention +models use the coalescer without this checkpoint wait. | Variable | Type | Default | Description | |----------|------|---------|-------------| @@ -37,8 +44,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. Also bounds local checkpoint waits. Empty disables the age guard. 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/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_scheduler_partial_prefill_tail.py b/tests/test_scheduler_partial_prefill_tail.py index aa744953d9..3f3183eeff 100644 --- a/tests/test_scheduler_partial_prefill_tail.py +++ b/tests/test_scheduler_partial_prefill_tail.py @@ -45,7 +45,8 @@ 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`.""" - dp_size = 2 + is_local = False + max_queue_ms = None def protects_decode(self, running_decode_batch: int) -> bool: return False From 23fdaa3ea8d91822b5fb0c8ee8ec640a8028b120 Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 22 Sep 2026 07:35:48 +0000 Subject: [PATCH 5/8] fix(scheduler): preserve checkpoint waits across queue bypass Keep per-request deadlines through retries and allocation rollbacks while allowing bounded lookahead for independent requests. Wait for useful in-flight checkpoints beyond smaller offload hits, without letting prior queue age disable prefix reuse. Clarify coalescing and checkpoint wait limits. --- atom/model_engine/scheduler.py | 53 +++++++++++++++++++--------------- atom/utils/envs.py | 8 ++--- docs/environment_variables.md | 17 +++++++---- tests/test_scheduler.py | 2 ++ 4 files changed, 46 insertions(+), 34 deletions(-) diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index abc8532e43..84d1ca825e 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -703,12 +703,12 @@ def __init__( self.prefill_delayer: PrefillDelayer | None = None self._local_prefill_coalescing = False - self._inflight_prefix_wait: tuple[int, int] | None = None + self._inflight_prefix_wait: dict[int, int] = {} def set_prefill_delayer(self, delayer) -> None: self.prefill_delayer = delayer self._local_prefill_coalescing = delayer is not None and delayer.is_local - self._inflight_prefix_wait = None + 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.""" @@ -722,14 +722,8 @@ def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: or seq.multimodal_data is not None ): return False - if self._inflight_prefix_wait is not None: - waiting_id, deadline = self._inflight_prefix_wait - if waiting_id == seq.id and self._schedule_tick >= deadline: - return False - if ( - delayer.max_queue_ms is not None - and (time.time() - seq.arrive_time) * 1000 >= delayer.max_queue_ms - ): + 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 @@ -762,14 +756,9 @@ def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: np.frombuffer(seq.token_ids, dtype=np.int32, count=anchor), np.frombuffer(producer.token_ids, dtype=np.int32, count=anchor), ): - if ( - self._inflight_prefix_wait is None - or self._inflight_prefix_wait[0] != seq.id - ): - self._inflight_prefix_wait = ( - seq.id, - self._schedule_tick + delayer.ttft_max_ticks, - ) + self._inflight_prefix_wait.setdefault( + seq.id, self._schedule_tick + delayer.ttft_max_ticks + ) return True return False @@ -1558,6 +1547,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())) @@ -1565,6 +1556,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 @@ -1595,6 +1590,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 @@ -1717,14 +1713,17 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: ) if ( self._local_prefill_coalescing - and not needs_remote_load - and not seq.offload_joint.kv_prefix_tokens + and (not needs_remote_load or self._connector_flag("is_offload")) and self._wait_for_inflight_prefix( - seq, num_cached_blocks * self.block_manager.hash_block_size + seq, + max( + num_cached_blocks * self.block_manager.hash_block_size, + seq.offload_joint.kv_prefix_tokens, + ), ) ): - self.waiting.appendleft(seq) - break + prefix_waiters.append(seq) + continue atomic_prefill = self._requires_atomic_prefill(seq) if ( not atomic_prefill @@ -1787,6 +1786,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: @@ -1823,6 +1823,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 @@ -1863,6 +1864,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( @@ -1870,6 +1872,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( @@ -2209,6 +2215,7 @@ 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" diff --git a/atom/utils/envs.py b/atom/utils/envs.py index 834e08147a..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", "") == "" diff --git a/docs/environment_variables.md b/docs/environment_variables.md index 115bc29e1a..4f2fa05e10 100644 --- a/docs/environment_variables.md +++ b/docs/environment_variables.md @@ -29,11 +29,16 @@ 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. -Local hybrid models with state checkpointing also wait briefly for an in-flight -producer's reusable prompt-end checkpoint. This preserves queue order, bypasses -requests with remote KV matches, and expires after `TTFT_MAX_TICKS` scheduler -passes or `MAX_QUEUE_MS` since arrival, whichever applies first. Pure-attention -models use the coalescer without this checkpoint wait. +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 | |----------|------|---------|-------------| @@ -44,7 +49,7 @@ models use the coalescer without this checkpoint wait. | **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) | After decode protection, release coalescing when the oldest schedulable waiting prefill reaches this age since arrival. Also bounds local checkpoint waits. Empty disables the age guard. This is not a hard TTFT limit. | +| **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/tests/test_scheduler.py b/tests/test_scheduler.py index cb8b4a1b03..d34c71274c 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -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 = [] @@ -508,6 +509,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 = [] From 09638d109c8a26ad93d2bebcd2c98addde79eb56 Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 22 Sep 2026 07:50:13 +0000 Subject: [PATCH 6/8] test(scheduler): initialize prefix waits in offload abort fixtures Three abort tests bypass Scheduler.__init__ and missed its new deadline map. Initialize the map so all six parameterized offload cancellation cases exercise resource cleanup with a complete scheduler fixture. --- tests/test_lmcache_offload_connector.py | 3 +++ 1 file changed, 3 insertions(+) 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, From 66e6bd01e80a5820c19fddbd31ab388c49ad68de Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 22 Sep 2026 08:25:37 +0000 Subject: [PATCH 7/8] fix(scheduler): align local prefill probes and reclaim cancellations --- atom/model_engine/block_manager.py | 52 +-- atom/model_engine/engine_core.py | 6 +- atom/model_engine/engine_utility.py | 17 +- atom/model_engine/prefill_delayer.py | 10 +- atom/model_engine/scheduler.py | 193 +++++++---- docs/environment_variables.md | 15 + docs/scheduling_kv_cache_guide.md | 81 ++--- tests/test_local_prefill.py | 480 +++++++++++++++++++++++++++ tests/test_scheduler.py | 7 +- 9 files changed, 725 insertions(+), 136 deletions(-) create mode 100644 tests/test_local_prefill.py diff --git a/atom/model_engine/block_manager.py b/atom/model_engine/block_manager.py index 2cb8603655..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) @@ -867,12 +873,14 @@ def can_allocate( 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` returns the same fit and HBM hit without committing - joint-load fields or funnel counters. It still refreshes checkpoint - demand/end and hit instrumentation. An optional `block_hashes` output + 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. @@ -906,9 +914,21 @@ def can_allocate( # match). Record each block's hash for the SWA scan below. h = seq.cache_seed compressed_hit = 0 + 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 @@ -955,22 +975,9 @@ def can_allocate( 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: self.record_allocation(seq, num_cached_blocks, block_hashes) # After the refusal, not before it. The chain is O(prompt) xxhash plus @@ -2073,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. @@ -2146,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: @@ -2620,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 b3851b7b6d..319ab8edce 100644 --- a/atom/model_engine/engine_core.py +++ b/atom/model_engine/engine_core.py @@ -195,7 +195,11 @@ def __init__(self, config: Config, input_address: str, output_address: str): 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: + 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 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 546923bb75..ca6c20776c 100644 --- a/atom/model_engine/prefill_delayer.py +++ b/atom/model_engine/prefill_delayer.py @@ -294,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. @@ -306,7 +309,8 @@ 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 age guard; all ranks release after decode diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index 84d1ca825e..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") @@ -699,13 +703,11 @@ def __init__( ) # Set by EngineCore for cross-DP alignment or opt-in TP decode protection. - from atom.model_engine.prefill_delayer import PrefillDelayer - 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 = {} @@ -739,6 +741,8 @@ def _wait_for_inflight_prefix(self, seq: Sequence, cached_tokens: int) -> bool: 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( @@ -772,39 +776,58 @@ def _waiting_prefills(self): yield seq def _local_prefill_pending_work(self) -> tuple[bool, int]: - """Estimate HBM-discounted work; cap further probes after one prompt budget.""" + """Count probed work; a bounded scan never invents fill for unseen work.""" budget = self.max_num_batched_tokens - pending = self._partial_prefill_remaining_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(): - if self._is_offload_prefill_resume(seq): - # Already owns its KV and state slots. Mirror the one-block - # recompute used by Phase 2 when the tier covers the whole seq. - cached = min( - seq.num_cached_tokens, - (seq.num_tokens - 1) - // self.block_manager.hash_block_size - * self.block_manager.hash_block_size, - ) + offload_resume = self._is_offload_prefill_resume(seq) + if offload_resume: + cached = self._offload_prefill_start(seq) else: if probed_tokens >= budget: - # A deep cache-heavy queue must not turn into a fixed stall - # delay just because estimating its fill is expensive. - return prefillable, budget if prefillable else pending - cached_blocks = self.block_manager.can_allocate(seq, record=False) + 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 not self.enable_chunked_prefill and remaining > budget: - continue + 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 += remaining + pending += chunk slots -= 1 if pending >= budget: return True, budget @@ -812,6 +835,40 @@ def _local_prefill_pending_work(self) -> tuple[bool, int]: 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 this rank would *actually* admit a new prefill this tick. @@ -911,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 @@ -979,12 +1036,10 @@ 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: @@ -1478,6 +1533,9 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: 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 @@ -1527,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 @@ -1648,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( @@ -1698,7 +1749,10 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: # where the trailing-window SWA is present); no post-hoc warmup trim. block_hashes = [] num_cached_blocks = self.block_manager.can_allocate( - seq, record=False, block_hashes=block_hashes + seq, + record=False, + block_hashes=block_hashes, + reuse_hashes=self._local_prefill_coalescing, ) if num_cached_blocks < 0: self.waiting.appendleft(seq) @@ -1708,8 +1762,8 @@ 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 ) if ( self._local_prefill_coalescing @@ -1725,12 +1779,6 @@ def _schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]] | None: prefix_waiters.append(seq) continue 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 - ): - num_new_tokens = self.long_prefill_token_threshold budget_remaining = self.max_num_batched_tokens - num_batched_tokens chunk = self._prefill_chunk_for_budget( num_new_tokens, budget_remaining, num_batched_tokens @@ -2221,20 +2269,12 @@ def _reject_aborted_waiting(self, seq: Sequence) -> None: 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 @@ -2250,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") @@ -3681,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 (): @@ -3694,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, @@ -3841,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 @@ -4114,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/docs/environment_variables.md b/docs/environment_variables.md index 4f2fa05e10..ae7e442d61 100644 --- a/docs/environment_variables.md +++ b/docs/environment_variables.md @@ -22,6 +22,8 @@ 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 @@ -29,6 +31,19 @@ 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 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_local_prefill.py b/tests/test_local_prefill.py new file mode 100644 index 0000000000..6cc59befab --- /dev/null +++ b/tests/test_local_prefill.py @@ -0,0 +1,480 @@ +# SPDX-License-Identifier: MIT +"""Local coalescing signals, checkpoint waits and cancellation lifetimes.""" + +import ast +import gc +import queue +import time +import weakref +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +from conftest import MockConfig, atom_config_double +from test_state_checkpoint import DEFAULT_STATE_RUNTIME + +from atom.kv_transfer.disaggregation.types import KVConnectorOutput +from atom.model_engine.engine_utility import EngineUtilityHandler +from atom.model_engine.prefill_delayer import PrefillDelayer +from atom.model_engine.scheduler import DecodeScheduler, ScheduledBatchOutput, Scheduler +from atom.model_engine.sequence import Sequence, SequenceStatus +from atom.model_engine.state_runtime import StateRuntime, StateTransfer +from atom.utils import envs + + +def scheduler(hybrid=False, **kwargs): + config = { + "max_num_seqs": 64, + "num_kvcache_blocks": 4096, + "kv_cache_block_size": 4, + "max_model_len": 65536, + "max_num_batched_tokens": 64, + "enable_prefix_caching": True, + "pool_entries": {"state": 64} if hybrid else {}, + "state_checkpoint_interval_tokens": 32, + } + config.update(kwargs) + sched = Scheduler( + MockConfig(**config), + **({"state_runtime": DEFAULT_STATE_RUNTIME} if hybrid else {}), + ) + sched.set_prefill_delayer( + PrefillDelayer(1, None, sched.max_num_batched_tokens, prefill_decode_interval=4) + ) + return sched + + +def sequence(n=128, hybrid=False, block_size=4): + seq = Sequence(list(range(n)), block_size, has_per_req_cache=hybrid) + seq.arrive_time = time.time() + return seq + + +def producer_waiter(): + sched = scheduler(hybrid=True, long_prefill_token_threshold=16) + producer = sequence(hybrid=True) + bm = sched.block_manager + assert bm.allocate(producer, bm.can_allocate(producer)) + producer.status = SequenceStatus.RUNNING + producer.num_cached_tokens = 16 + producer.is_partial_prefill = True + sched.running.append(producer) + sched._partial_prefill_count = 1 + return sched, producer, sequence(hybrid=True) + + +def publish(sched, seq): + bm = sched.block_manager + assert bm.allocate(seq, bm.can_allocate(seq)) + bm.hash_blocks(seq, seq.num_tokens) + bm.deallocate(seq) + + +def test_high_hit_probe_cap_does_not_invent_fill(): + sched = scheduler(max_num_batched_tokens=64) + publish(sched, sequence()) + sched.waiting.extend(sequence() for _ in range(8)) + assert sched._local_prefill_pending_work() == (True, 4) + d = sched.prefill_delayer + d._first = False + d.stall_ticks = 2 + for _ in range(3): + ready, pending = sched._local_prefill_pending_work() + released = d.should_allow_prefill(ready, pending, running_decode_batch=1) + assert released + assert d._stat_fire_fill == 0 + assert d._stat_fire_stall == 1 + + +def test_short_burst_can_fill_past_four_requests(): + sched = scheduler() + sched.waiting.extend(sequence(8) for _ in range(8)) + assert sched._local_prefill_pending_work() == (True, 64) + batch, _ = sched.schedule() + assert batch.total_tokens_num == 64 + + +@pytest.mark.parametrize("threshold,expected", [(0, 64), (16, 16)]) +def test_new_chunk_estimate_matches_admission(threshold, expected): + sched = scheduler(long_prefill_token_threshold=threshold) + sched.waiting.append(sequence()) + assert sched._local_prefill_pending_work() == (True, expected) + batch, _ = sched.schedule() + assert batch.total_tokens_num == expected + + +def test_nonchunked_refusal_stops_estimate_and_admission(): + sched = scheduler(enable_chunked_prefill=False) + first = sequence(40) + preempted = sequence(32) + for _ in range(8): + preempted.append_token(500) + tail = sequence(24) + sched.waiting.extend([first, preempted, tail]) + assert sched._local_prefill_pending_work() == (True, 40) + batch, admitted = sched.schedule() + assert batch.total_tokens_num == 40 + assert list(admitted) == [first.id] + assert list(sched.waiting) == [preempted, tail] + + +@pytest.mark.parametrize("cached,expected", [(900, 100), (1000, 232)]) +def test_offload_resume_estimate_matches_admission(cached, expected, monkeypatch): + sched = scheduler(kv_cache_block_size=256, max_num_batched_tokens=1024) + seq = sequence(1000, block_size=256) + assert sched.block_manager.allocate(seq, 0) + seq.num_cached_tokens = cached + sched.waiting.append(seq) + monkeypatch.setattr(sched, "_is_offload_prefill_resume", lambda value: value is seq) + assert sched._local_prefill_pending_work() == (True, expected) + batch, _ = sched.schedule() + assert batch.total_tokens_num == expected + assert seq.prefix_cache_hit_tokens == cached + + +def test_parked_remote_slots_do_not_report_fill(monkeypatch): + sched = scheduler(max_num_seqs=1) + sched._num_parked_remote_kv = 1 + sched.waiting.append(sequence(8)) + monkeypatch.setattr(sched, "_query_connector_prefill_match", lambda *a, **kw: True) + assert sched._local_prefill_pending_work() == (True, 0) + _, admitted = sched.schedule() + assert not admitted + + +def test_partial_preview_preserves_checkpoint_cut_counter(): + sched, producer, _ = producer_waiter() + sched.long_prefill_token_threshold = 64 + producer.checkpoint_demand_pos = 36 + assert sched._local_prefill_pending_work() == (True, 20) + assert sched.block_manager.chunks_cut_for_demand == 0 + batch, _ = sched.schedule() + assert batch.total_tokens_num == 20 + assert sched.block_manager.chunks_cut_for_demand == 1 + + +def test_each_waiter_keeps_its_deadline_after_head_changes(): + sched, _, first = producer_waiter() + second = sequence(hybrid=True) + sched.prefill_delayer.ttft_max_ticks = 3 + assert sched._wait_for_inflight_prefix(first, 0) + sched._schedule_tick += 1 + assert sched._wait_for_inflight_prefix(second, 0) + deadline = sched._inflight_prefix_wait[first.id] + sched._schedule_tick = deadline + assert not sched._wait_for_inflight_prefix(first, 0) + assert sched._wait_for_inflight_prefix(second, 0) + assert sched._inflight_prefix_wait[first.id] == deadline + + +def test_waiting_head_allows_independent_request(): + sched, _, waiter = producer_waiter() + independent = sequence(8, hybrid=True) + independent.token_ids[0] = 90000 + sched.waiting.extend([waiter, independent]) + _, admitted = sched.schedule() + assert independent.id in admitted + assert waiter.id not in admitted + assert not waiter.block_table + assert list(sched.waiting) == [waiter] + + +def test_bypass_scan_is_bounded_and_preserves_waiters(monkeypatch): + sched, _, waiter = producer_waiter() + sched.waiting.extend([waiter, *[sequence(hybrid=True) for _ in range(31)]]) + original = list(sched.waiting) + joint_skips = dict(sched.block_manager.joint_skips) + probe = Mock(wraps=sched.block_manager.can_allocate) + monkeypatch.setattr(sched.block_manager, "can_allocate", probe) + sched.schedule() + assert list(sched.waiting) == original + assert probe.call_count == 17 + assert len(sched._inflight_prefix_wait) == 17 + assert sched.block_manager.joint_skips == joint_skips + + +def test_failed_allocation_keeps_expired_deadline(monkeypatch): + sched, _, waiter = producer_waiter() + assert sched._wait_for_inflight_prefix(waiter, 0) + deadline = sched._inflight_prefix_wait[waiter.id] + sched._schedule_tick = deadline + sched.waiting.append(waiter) + monkeypatch.setattr(sched.block_manager, "allocate", lambda *args: False) + sched.schedule() + assert sched._inflight_prefix_wait[waiter.id] == deadline + assert not sched._wait_for_inflight_prefix(waiter, 0) + + +def test_waiter_reuses_completed_prefix(): + sched = scheduler( + hybrid=True, max_num_batched_tokens=16, state_checkpoint_interval_tokens=64 + ) + producer = sequence(60, hybrid=True) + waiter = sequence(76, hybrid=True) + sched.extend([producer, waiter]) + was_deferred = False + for _ in range(40): + batch, admitted = sched.schedule() + if waiter.id in admitted: + assert was_deferred + assert waiter.prefix_cache_hit_tokens >= 48 + assert waiter.id not in sched._inflight_prefix_wait + return + was_deferred |= waiter.id in sched._inflight_prefix_wait + assert not waiter.block_table + sched.postprocess( + list(admitted.values()), + ScheduledBatchOutput( + req_ids=batch.req_ids, + token_ids=[(501,)] * len(batch.req_ids), + num_rejected=None, + num_bonus=None, + draft_token_ids=None, + ), + batch=batch, + ) + if batch.total_seqs_num_prefill: + sched.prefill_delayer.notify_prefill_executed() + pytest.fail("waiter never admitted") + + +def test_consumer_must_have_room_to_restore_fork(): + sched = scheduler(hybrid=True, max_num_batched_tokens=32) + sched.block_manager = type(sched.block_manager)( + MockConfig( + kv_cache_block_size=64, + num_kvcache_blocks=100, + pool_entries={"state": 8}, + enable_prefix_caching=True, + state_checkpoint_interval_tokens=128, + ), + state_runtime=StateRuntime(transfer=StateTransfer.fork(131)), + ) + producer = sequence(1024, hybrid=True, block_size=64) + producer.status = SequenceStatus.RUNNING + producer.is_partial_prefill = True + producer.checkpoint_end_pos = 832 + sched.running.append(producer) + short = sequence(896, hybrid=True, block_size=64) + long = sequence(1024, hybrid=True, block_size=64) + assert not sched._wait_for_inflight_prefix(short, 0) + assert sched._wait_for_inflight_prefix(long, 0) + pool = sched.block_manager.state + pool._index(13, 0) + assert pool.resumable_hit(short, 13, list(range(1, 14))) == 0 + assert pool.resumable_hit(long, 13, list(range(1, 14))) == 13 + + +@pytest.mark.parametrize("dcp", [1, 8]) +def test_prompt_hash_reuse_rechecks_pool_and_seed(dcp, monkeypatch): + sched = scheduler(decode_context_parallel_size=dcp) + bm = sched.block_manager + if dcp > 1: + # Hash/pool logic is real; only the GPU module's shard sizing is isolated. + monkeypatch.setattr( + bm, "num_pool_blocks", lambda n: (n + 4 * dcp - 1) // (4 * dcp) + ) + publish(sched, sequence()) + seq = sequence() + hashes = Mock(wraps=bm.compute_hash) + monkeypatch.setattr(bm, "compute_hash", hashes) + expected = 128 // (4 * dcp) - 1 + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected + first_calls = hashes.call_count + assert first_calls > 0 + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected + assert hashes.call_count == first_calls + original_seed = seq.cache_seed + seq.cache_seed = 123 + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 0 + seq.cache_seed = original_seed + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected + monkeypatch.setattr(bm.kv, "lookup", lambda h: -1) + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 0 + + +def test_generated_suffix_is_not_memoized_and_weak_cache_releases_request(): + sched = scheduler() + seq = sequence(10) + for token in range(10, 24): + seq.append_token(token) + publish(sched, seq) + bm = sched.block_manager + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 5 + assert len(bm._prefill_probe_hashes[seq][1]) == 2 + seq.token_ids[12] = 10000 + assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 3 + ref = weakref.ref(seq) + del seq + gc.collect() + assert ref() is None + assert not bm._prefill_probe_hashes + + +def test_deallocation_clears_old_joint_span(): + sched = scheduler() + seq = sequence(8) + bm = sched.block_manager + assert bm.allocate(seq, 0) + seq.offload_joint.boundary_tokens = 4 + seq.offload_joint.boundary_hash = 12 + bm.deallocate(seq) + assert not seq.block_table + assert seq.offload_joint.boundary_tokens == 0 + assert seq.offload_joint.boundary_hash == -1 + + +def test_nonhead_abort_reclaims_slot_before_admission(): + sched = scheduler(hybrid=True, pool_entries={"state": 1}) + head = sequence(8, hybrid=True) + aborted = sequence(8, hybrid=True) + assert sched.block_manager.allocate(aborted, 0) + sched.waiting.extend([head, aborted]) + assert sched.block_manager.can_allocate(head, record=False) == -1 + handler = EngineUtilityHandler(None, queue.Queue(), scheduler=sched) + handler._handle_abort_request({"req_id": aborted.id}) + assert not aborted.block_table and not aborted.state_slots + assert aborted in sched.take_rejected() + _, admitted = sched.schedule() + assert head.id in admitted + + +@pytest.mark.parametrize( + "terminal", + ["finished_recving", "failed_recving", "finished_loading", "failed_loading"], +) +def test_nonhead_abort_retains_inflight_resources_until_terminal(terminal): + sched = scheduler(hybrid=True, pool_entries={"state": 1}) + head, aborted = sequence(8, hybrid=True), sequence(8, hybrid=True) + assert sched.block_manager.allocate(aborted, 0) + sched.waiting.extend([head, aborted]) + sched._count_inflight_load(aborted) + aborted.status = SequenceStatus.WAITING_FOR_REMOTE_KVS + sched.kv_connector = SimpleNamespace( + is_producer=False, is_offload="loading" in terminal + ) + assert sched.abort_request(aborted.id) + assert aborted.block_table and aborted.state_slots + assert sched._num_parked_remote_kv == 1 + sched._update_from_kv_xfer_finished(KVConnectorOutput(**{terminal: {aborted.id}})) + assert not aborted.block_table and not aborted.state_slots + assert sched._num_parked_remote_kv == 0 + assert not sched.deferred_free_blocks + + +def test_running_abort_keeps_forward_resources(): + sched = scheduler() + seq = sequence(8) + sched.add(seq) + sched.schedule() + assert sched.abort_request(seq.id) + assert seq.status == SequenceStatus.ABORTED + assert seq.block_table + assert seq in sched.running + assert not sched.abort_request(-1) + + +def test_protection_skips_probes_then_resumes_partial(monkeypatch): + sched, producer, _ = producer_waiter() + decode = sequence(8, hybrid=True) + assert sched.block_manager.allocate(decode, 0) + decode.num_cached_tokens = decode.num_prompt_tokens + decode.append_token(500) + decode.status = SequenceStatus.RUNNING + sched.running.appendleft(decode) + d = sched.prefill_delayer + d._first = False + d.partial_max_ticks = 0 + d.notify_prefill_executed() + with monkeypatch.context() as patch: + patch.setattr( + sched, "_local_prefill_pending_work", Mock(side_effect=AssertionError) + ) + for _ in range(4): + batch, admitted = sched.schedule() + assert producer.id not in admitted + assert batch.total_seqs_num_prefill == 0 + assert d._hold_ticks == 0 + batch, admitted = sched.schedule() + assert producer.id in admitted + assert batch.total_seqs_num_prefill == 1 + assert d._stat_fire_partial == 1 + + +@pytest.mark.parametrize( + "rapidserve,has_scheduler,expected", + [(True, True, False), (False, False, False), (False, True, True)], +) +def test_delayer_init_distinguishes_rapidserve_from_connector_pd( + rapidserve, has_scheduler, expected, monkeypatch +): + monkeypatch.setenv("ATOM_ENABLE_PREFILL_DELAYER", "1") + config = atom_config_double(enable_rapidserve=rapidserve, max_num_batched_tokens=64) + config.parallel_config = SimpleNamespace(data_parallel_size=1) + # Execute the real helper without importing AITER's worker IPC transport. + source = Path(__file__).resolve().parents[1] / "atom/model_engine/engine_core.py" + tree = ast.parse(source.read_text()) + cls = next( + n for n in tree.body if isinstance(n, ast.ClassDef) and n.name == "EngineCore" + ) + method = next( + n + for n in cls.body + if isinstance(n, ast.FunctionDef) and n.name == "_init_prefill_delayer" + ) + namespace = {"Config": object, "envs": envs} + exec( # noqa: S102 — execute the repository helper, without GPU-only imports. + compile(ast.Module(body=[method], type_ignores=[]), str(source), "exec"), + namespace, + ) + core = SimpleNamespace() + core.scheduler = scheduler() if has_scheduler else None + if core.scheduler: + core.scheduler.set_prefill_delayer(None) + core.scheduler.kv_connector = SimpleNamespace(is_producer=False) + namespace["_init_prefill_delayer"](core, config) + attached = core.scheduler is not None and core.scheduler.prefill_delayer is not None + assert attached == expected + + +def test_rapidserve_decode_abort_retains_its_mark_only_path(): + sched = DecodeScheduler(MockConfig()) + seq = sequence(8) + sched.waiting.append(seq) + assert sched.abort_request(seq.id) + assert seq.status == SequenceStatus.ABORTED + assert seq in sched.waiting + assert not sched._rejected + + +def test_nonhead_completed_offload_abort_reclaims_slot(): + sched = scheduler(hybrid=True, pool_entries={"state": 1}) + head, aborted = sequence(8, hybrid=True), sequence(8, hybrid=True) + assert sched.block_manager.allocate(aborted, 0) + sched.waiting.extend([head, aborted]) + sched._count_inflight_load(aborted) + aborted.offload_loaded = True + sched.kv_connector = SimpleNamespace(is_producer=False, is_offload=True) + assert sched.block_manager.can_allocate(head, record=False) == -1 + assert sched.abort_request(aborted.id) + assert not aborted.block_table and not aborted.state_slots + assert sched._num_parked_remote_kv == 0 + assert sched.block_manager.can_allocate(head, record=False) >= 0 + + +def test_parked_slots_still_allow_local_work_after_stall(): + sched = scheduler(max_num_seqs=1) + sched._num_parked_remote_kv = 1 + seq = sequence(8) + sched.waiting.append(seq) + d = sched.prefill_delayer + d._first = False + d.stall_ticks = 1 + ready, pending = sched._local_prefill_pending_work() + assert (ready, pending) == (True, 0) + assert not d.should_allow_prefill(ready, pending, running_decode_batch=1) + assert d.should_allow_prefill(ready, pending, running_decode_batch=1) + assert d._stat_fire_stall == 1 and d._stat_fire_vacuous == 0 + _, admitted = sched.schedule() + assert seq.id in admitted diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index d34c71274c..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, @@ -445,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]) From 1e7af70e02f7ea8c773038571a68bce48a98c167 Mon Sep 17 00:00:00 2001 From: whx-sjtu Date: Tue, 22 Sep 2026 09:14:16 +0000 Subject: [PATCH 8/8] test: remove local prefill regression module --- tests/test_local_prefill.py | 480 ------------------------------------ 1 file changed, 480 deletions(-) delete mode 100644 tests/test_local_prefill.py diff --git a/tests/test_local_prefill.py b/tests/test_local_prefill.py deleted file mode 100644 index 6cc59befab..0000000000 --- a/tests/test_local_prefill.py +++ /dev/null @@ -1,480 +0,0 @@ -# SPDX-License-Identifier: MIT -"""Local coalescing signals, checkpoint waits and cancellation lifetimes.""" - -import ast -import gc -import queue -import time -import weakref -from pathlib import Path -from types import SimpleNamespace -from unittest.mock import Mock - -import pytest -from conftest import MockConfig, atom_config_double -from test_state_checkpoint import DEFAULT_STATE_RUNTIME - -from atom.kv_transfer.disaggregation.types import KVConnectorOutput -from atom.model_engine.engine_utility import EngineUtilityHandler -from atom.model_engine.prefill_delayer import PrefillDelayer -from atom.model_engine.scheduler import DecodeScheduler, ScheduledBatchOutput, Scheduler -from atom.model_engine.sequence import Sequence, SequenceStatus -from atom.model_engine.state_runtime import StateRuntime, StateTransfer -from atom.utils import envs - - -def scheduler(hybrid=False, **kwargs): - config = { - "max_num_seqs": 64, - "num_kvcache_blocks": 4096, - "kv_cache_block_size": 4, - "max_model_len": 65536, - "max_num_batched_tokens": 64, - "enable_prefix_caching": True, - "pool_entries": {"state": 64} if hybrid else {}, - "state_checkpoint_interval_tokens": 32, - } - config.update(kwargs) - sched = Scheduler( - MockConfig(**config), - **({"state_runtime": DEFAULT_STATE_RUNTIME} if hybrid else {}), - ) - sched.set_prefill_delayer( - PrefillDelayer(1, None, sched.max_num_batched_tokens, prefill_decode_interval=4) - ) - return sched - - -def sequence(n=128, hybrid=False, block_size=4): - seq = Sequence(list(range(n)), block_size, has_per_req_cache=hybrid) - seq.arrive_time = time.time() - return seq - - -def producer_waiter(): - sched = scheduler(hybrid=True, long_prefill_token_threshold=16) - producer = sequence(hybrid=True) - bm = sched.block_manager - assert bm.allocate(producer, bm.can_allocate(producer)) - producer.status = SequenceStatus.RUNNING - producer.num_cached_tokens = 16 - producer.is_partial_prefill = True - sched.running.append(producer) - sched._partial_prefill_count = 1 - return sched, producer, sequence(hybrid=True) - - -def publish(sched, seq): - bm = sched.block_manager - assert bm.allocate(seq, bm.can_allocate(seq)) - bm.hash_blocks(seq, seq.num_tokens) - bm.deallocate(seq) - - -def test_high_hit_probe_cap_does_not_invent_fill(): - sched = scheduler(max_num_batched_tokens=64) - publish(sched, sequence()) - sched.waiting.extend(sequence() for _ in range(8)) - assert sched._local_prefill_pending_work() == (True, 4) - d = sched.prefill_delayer - d._first = False - d.stall_ticks = 2 - for _ in range(3): - ready, pending = sched._local_prefill_pending_work() - released = d.should_allow_prefill(ready, pending, running_decode_batch=1) - assert released - assert d._stat_fire_fill == 0 - assert d._stat_fire_stall == 1 - - -def test_short_burst_can_fill_past_four_requests(): - sched = scheduler() - sched.waiting.extend(sequence(8) for _ in range(8)) - assert sched._local_prefill_pending_work() == (True, 64) - batch, _ = sched.schedule() - assert batch.total_tokens_num == 64 - - -@pytest.mark.parametrize("threshold,expected", [(0, 64), (16, 16)]) -def test_new_chunk_estimate_matches_admission(threshold, expected): - sched = scheduler(long_prefill_token_threshold=threshold) - sched.waiting.append(sequence()) - assert sched._local_prefill_pending_work() == (True, expected) - batch, _ = sched.schedule() - assert batch.total_tokens_num == expected - - -def test_nonchunked_refusal_stops_estimate_and_admission(): - sched = scheduler(enable_chunked_prefill=False) - first = sequence(40) - preempted = sequence(32) - for _ in range(8): - preempted.append_token(500) - tail = sequence(24) - sched.waiting.extend([first, preempted, tail]) - assert sched._local_prefill_pending_work() == (True, 40) - batch, admitted = sched.schedule() - assert batch.total_tokens_num == 40 - assert list(admitted) == [first.id] - assert list(sched.waiting) == [preempted, tail] - - -@pytest.mark.parametrize("cached,expected", [(900, 100), (1000, 232)]) -def test_offload_resume_estimate_matches_admission(cached, expected, monkeypatch): - sched = scheduler(kv_cache_block_size=256, max_num_batched_tokens=1024) - seq = sequence(1000, block_size=256) - assert sched.block_manager.allocate(seq, 0) - seq.num_cached_tokens = cached - sched.waiting.append(seq) - monkeypatch.setattr(sched, "_is_offload_prefill_resume", lambda value: value is seq) - assert sched._local_prefill_pending_work() == (True, expected) - batch, _ = sched.schedule() - assert batch.total_tokens_num == expected - assert seq.prefix_cache_hit_tokens == cached - - -def test_parked_remote_slots_do_not_report_fill(monkeypatch): - sched = scheduler(max_num_seqs=1) - sched._num_parked_remote_kv = 1 - sched.waiting.append(sequence(8)) - monkeypatch.setattr(sched, "_query_connector_prefill_match", lambda *a, **kw: True) - assert sched._local_prefill_pending_work() == (True, 0) - _, admitted = sched.schedule() - assert not admitted - - -def test_partial_preview_preserves_checkpoint_cut_counter(): - sched, producer, _ = producer_waiter() - sched.long_prefill_token_threshold = 64 - producer.checkpoint_demand_pos = 36 - assert sched._local_prefill_pending_work() == (True, 20) - assert sched.block_manager.chunks_cut_for_demand == 0 - batch, _ = sched.schedule() - assert batch.total_tokens_num == 20 - assert sched.block_manager.chunks_cut_for_demand == 1 - - -def test_each_waiter_keeps_its_deadline_after_head_changes(): - sched, _, first = producer_waiter() - second = sequence(hybrid=True) - sched.prefill_delayer.ttft_max_ticks = 3 - assert sched._wait_for_inflight_prefix(first, 0) - sched._schedule_tick += 1 - assert sched._wait_for_inflight_prefix(second, 0) - deadline = sched._inflight_prefix_wait[first.id] - sched._schedule_tick = deadline - assert not sched._wait_for_inflight_prefix(first, 0) - assert sched._wait_for_inflight_prefix(second, 0) - assert sched._inflight_prefix_wait[first.id] == deadline - - -def test_waiting_head_allows_independent_request(): - sched, _, waiter = producer_waiter() - independent = sequence(8, hybrid=True) - independent.token_ids[0] = 90000 - sched.waiting.extend([waiter, independent]) - _, admitted = sched.schedule() - assert independent.id in admitted - assert waiter.id not in admitted - assert not waiter.block_table - assert list(sched.waiting) == [waiter] - - -def test_bypass_scan_is_bounded_and_preserves_waiters(monkeypatch): - sched, _, waiter = producer_waiter() - sched.waiting.extend([waiter, *[sequence(hybrid=True) for _ in range(31)]]) - original = list(sched.waiting) - joint_skips = dict(sched.block_manager.joint_skips) - probe = Mock(wraps=sched.block_manager.can_allocate) - monkeypatch.setattr(sched.block_manager, "can_allocate", probe) - sched.schedule() - assert list(sched.waiting) == original - assert probe.call_count == 17 - assert len(sched._inflight_prefix_wait) == 17 - assert sched.block_manager.joint_skips == joint_skips - - -def test_failed_allocation_keeps_expired_deadline(monkeypatch): - sched, _, waiter = producer_waiter() - assert sched._wait_for_inflight_prefix(waiter, 0) - deadline = sched._inflight_prefix_wait[waiter.id] - sched._schedule_tick = deadline - sched.waiting.append(waiter) - monkeypatch.setattr(sched.block_manager, "allocate", lambda *args: False) - sched.schedule() - assert sched._inflight_prefix_wait[waiter.id] == deadline - assert not sched._wait_for_inflight_prefix(waiter, 0) - - -def test_waiter_reuses_completed_prefix(): - sched = scheduler( - hybrid=True, max_num_batched_tokens=16, state_checkpoint_interval_tokens=64 - ) - producer = sequence(60, hybrid=True) - waiter = sequence(76, hybrid=True) - sched.extend([producer, waiter]) - was_deferred = False - for _ in range(40): - batch, admitted = sched.schedule() - if waiter.id in admitted: - assert was_deferred - assert waiter.prefix_cache_hit_tokens >= 48 - assert waiter.id not in sched._inflight_prefix_wait - return - was_deferred |= waiter.id in sched._inflight_prefix_wait - assert not waiter.block_table - sched.postprocess( - list(admitted.values()), - ScheduledBatchOutput( - req_ids=batch.req_ids, - token_ids=[(501,)] * len(batch.req_ids), - num_rejected=None, - num_bonus=None, - draft_token_ids=None, - ), - batch=batch, - ) - if batch.total_seqs_num_prefill: - sched.prefill_delayer.notify_prefill_executed() - pytest.fail("waiter never admitted") - - -def test_consumer_must_have_room_to_restore_fork(): - sched = scheduler(hybrid=True, max_num_batched_tokens=32) - sched.block_manager = type(sched.block_manager)( - MockConfig( - kv_cache_block_size=64, - num_kvcache_blocks=100, - pool_entries={"state": 8}, - enable_prefix_caching=True, - state_checkpoint_interval_tokens=128, - ), - state_runtime=StateRuntime(transfer=StateTransfer.fork(131)), - ) - producer = sequence(1024, hybrid=True, block_size=64) - producer.status = SequenceStatus.RUNNING - producer.is_partial_prefill = True - producer.checkpoint_end_pos = 832 - sched.running.append(producer) - short = sequence(896, hybrid=True, block_size=64) - long = sequence(1024, hybrid=True, block_size=64) - assert not sched._wait_for_inflight_prefix(short, 0) - assert sched._wait_for_inflight_prefix(long, 0) - pool = sched.block_manager.state - pool._index(13, 0) - assert pool.resumable_hit(short, 13, list(range(1, 14))) == 0 - assert pool.resumable_hit(long, 13, list(range(1, 14))) == 13 - - -@pytest.mark.parametrize("dcp", [1, 8]) -def test_prompt_hash_reuse_rechecks_pool_and_seed(dcp, monkeypatch): - sched = scheduler(decode_context_parallel_size=dcp) - bm = sched.block_manager - if dcp > 1: - # Hash/pool logic is real; only the GPU module's shard sizing is isolated. - monkeypatch.setattr( - bm, "num_pool_blocks", lambda n: (n + 4 * dcp - 1) // (4 * dcp) - ) - publish(sched, sequence()) - seq = sequence() - hashes = Mock(wraps=bm.compute_hash) - monkeypatch.setattr(bm, "compute_hash", hashes) - expected = 128 // (4 * dcp) - 1 - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected - first_calls = hashes.call_count - assert first_calls > 0 - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected - assert hashes.call_count == first_calls - original_seed = seq.cache_seed - seq.cache_seed = 123 - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 0 - seq.cache_seed = original_seed - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == expected - monkeypatch.setattr(bm.kv, "lookup", lambda h: -1) - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 0 - - -def test_generated_suffix_is_not_memoized_and_weak_cache_releases_request(): - sched = scheduler() - seq = sequence(10) - for token in range(10, 24): - seq.append_token(token) - publish(sched, seq) - bm = sched.block_manager - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 5 - assert len(bm._prefill_probe_hashes[seq][1]) == 2 - seq.token_ids[12] = 10000 - assert bm.can_allocate(seq, record=False, reuse_hashes=True) == 3 - ref = weakref.ref(seq) - del seq - gc.collect() - assert ref() is None - assert not bm._prefill_probe_hashes - - -def test_deallocation_clears_old_joint_span(): - sched = scheduler() - seq = sequence(8) - bm = sched.block_manager - assert bm.allocate(seq, 0) - seq.offload_joint.boundary_tokens = 4 - seq.offload_joint.boundary_hash = 12 - bm.deallocate(seq) - assert not seq.block_table - assert seq.offload_joint.boundary_tokens == 0 - assert seq.offload_joint.boundary_hash == -1 - - -def test_nonhead_abort_reclaims_slot_before_admission(): - sched = scheduler(hybrid=True, pool_entries={"state": 1}) - head = sequence(8, hybrid=True) - aborted = sequence(8, hybrid=True) - assert sched.block_manager.allocate(aborted, 0) - sched.waiting.extend([head, aborted]) - assert sched.block_manager.can_allocate(head, record=False) == -1 - handler = EngineUtilityHandler(None, queue.Queue(), scheduler=sched) - handler._handle_abort_request({"req_id": aborted.id}) - assert not aborted.block_table and not aborted.state_slots - assert aborted in sched.take_rejected() - _, admitted = sched.schedule() - assert head.id in admitted - - -@pytest.mark.parametrize( - "terminal", - ["finished_recving", "failed_recving", "finished_loading", "failed_loading"], -) -def test_nonhead_abort_retains_inflight_resources_until_terminal(terminal): - sched = scheduler(hybrid=True, pool_entries={"state": 1}) - head, aborted = sequence(8, hybrid=True), sequence(8, hybrid=True) - assert sched.block_manager.allocate(aborted, 0) - sched.waiting.extend([head, aborted]) - sched._count_inflight_load(aborted) - aborted.status = SequenceStatus.WAITING_FOR_REMOTE_KVS - sched.kv_connector = SimpleNamespace( - is_producer=False, is_offload="loading" in terminal - ) - assert sched.abort_request(aborted.id) - assert aborted.block_table and aborted.state_slots - assert sched._num_parked_remote_kv == 1 - sched._update_from_kv_xfer_finished(KVConnectorOutput(**{terminal: {aborted.id}})) - assert not aborted.block_table and not aborted.state_slots - assert sched._num_parked_remote_kv == 0 - assert not sched.deferred_free_blocks - - -def test_running_abort_keeps_forward_resources(): - sched = scheduler() - seq = sequence(8) - sched.add(seq) - sched.schedule() - assert sched.abort_request(seq.id) - assert seq.status == SequenceStatus.ABORTED - assert seq.block_table - assert seq in sched.running - assert not sched.abort_request(-1) - - -def test_protection_skips_probes_then_resumes_partial(monkeypatch): - sched, producer, _ = producer_waiter() - decode = sequence(8, hybrid=True) - assert sched.block_manager.allocate(decode, 0) - decode.num_cached_tokens = decode.num_prompt_tokens - decode.append_token(500) - decode.status = SequenceStatus.RUNNING - sched.running.appendleft(decode) - d = sched.prefill_delayer - d._first = False - d.partial_max_ticks = 0 - d.notify_prefill_executed() - with monkeypatch.context() as patch: - patch.setattr( - sched, "_local_prefill_pending_work", Mock(side_effect=AssertionError) - ) - for _ in range(4): - batch, admitted = sched.schedule() - assert producer.id not in admitted - assert batch.total_seqs_num_prefill == 0 - assert d._hold_ticks == 0 - batch, admitted = sched.schedule() - assert producer.id in admitted - assert batch.total_seqs_num_prefill == 1 - assert d._stat_fire_partial == 1 - - -@pytest.mark.parametrize( - "rapidserve,has_scheduler,expected", - [(True, True, False), (False, False, False), (False, True, True)], -) -def test_delayer_init_distinguishes_rapidserve_from_connector_pd( - rapidserve, has_scheduler, expected, monkeypatch -): - monkeypatch.setenv("ATOM_ENABLE_PREFILL_DELAYER", "1") - config = atom_config_double(enable_rapidserve=rapidserve, max_num_batched_tokens=64) - config.parallel_config = SimpleNamespace(data_parallel_size=1) - # Execute the real helper without importing AITER's worker IPC transport. - source = Path(__file__).resolve().parents[1] / "atom/model_engine/engine_core.py" - tree = ast.parse(source.read_text()) - cls = next( - n for n in tree.body if isinstance(n, ast.ClassDef) and n.name == "EngineCore" - ) - method = next( - n - for n in cls.body - if isinstance(n, ast.FunctionDef) and n.name == "_init_prefill_delayer" - ) - namespace = {"Config": object, "envs": envs} - exec( # noqa: S102 — execute the repository helper, without GPU-only imports. - compile(ast.Module(body=[method], type_ignores=[]), str(source), "exec"), - namespace, - ) - core = SimpleNamespace() - core.scheduler = scheduler() if has_scheduler else None - if core.scheduler: - core.scheduler.set_prefill_delayer(None) - core.scheduler.kv_connector = SimpleNamespace(is_producer=False) - namespace["_init_prefill_delayer"](core, config) - attached = core.scheduler is not None and core.scheduler.prefill_delayer is not None - assert attached == expected - - -def test_rapidserve_decode_abort_retains_its_mark_only_path(): - sched = DecodeScheduler(MockConfig()) - seq = sequence(8) - sched.waiting.append(seq) - assert sched.abort_request(seq.id) - assert seq.status == SequenceStatus.ABORTED - assert seq in sched.waiting - assert not sched._rejected - - -def test_nonhead_completed_offload_abort_reclaims_slot(): - sched = scheduler(hybrid=True, pool_entries={"state": 1}) - head, aborted = sequence(8, hybrid=True), sequence(8, hybrid=True) - assert sched.block_manager.allocate(aborted, 0) - sched.waiting.extend([head, aborted]) - sched._count_inflight_load(aborted) - aborted.offload_loaded = True - sched.kv_connector = SimpleNamespace(is_producer=False, is_offload=True) - assert sched.block_manager.can_allocate(head, record=False) == -1 - assert sched.abort_request(aborted.id) - assert not aborted.block_table and not aborted.state_slots - assert sched._num_parked_remote_kv == 0 - assert sched.block_manager.can_allocate(head, record=False) >= 0 - - -def test_parked_slots_still_allow_local_work_after_stall(): - sched = scheduler(max_num_seqs=1) - sched._num_parked_remote_kv = 1 - seq = sequence(8) - sched.waiting.append(seq) - d = sched.prefill_delayer - d._first = False - d.stall_ticks = 1 - ready, pending = sched._local_prefill_pending_work() - assert (ready, pending) == (True, 0) - assert not d.should_allow_prefill(ready, pending, running_decode_batch=1) - assert d.should_allow_prefill(ready, pending, running_decode_batch=1) - assert d._stat_fire_stall == 1 and d._stat_fire_vacuous == 0 - _, admitted = sched.schedule() - assert seq.id in admitted