diff --git a/atom/model_engine/block_manager.py b/atom/model_engine/block_manager.py index d9a9d47c2f..bff946aece 100644 --- a/atom/model_engine/block_manager.py +++ b/atom/model_engine/block_manager.py @@ -18,9 +18,16 @@ ) from atom.model_engine.block_pool import BlockPool from atom.model_engine.kv_block import STATE_SLOT_CLASS +from atom.model_engine.page_unit_checkpoint import PagedStateCheckpointCoordinator from atom.model_engine.sequence import Sequence -from atom.model_engine.state_cache import StateCache -from atom.model_engine.state_pool import StateGroupPool, StateTransfer +from atom.model_engine.state_cache import StateCache, StateCheckpointCache +from atom.model_engine.state_pool import StateGroupPool +from atom.model_engine.state_runtime import ( + DEFAULT_STATE_RUNTIME, + StateMaintenanceOps, + StateRuntime, + StateTransfer, +) logger = logging.getLogger("atom") @@ -51,7 +58,12 @@ def _make_all_cleared() -> AllBlocksCleared: class BlockManager: - def __init__(self, config: Config): + def __init__( + self, + config: Config, + *, + state_runtime: StateRuntime = DEFAULT_STATE_RUNTIME, + ): block_size = config.kv_cache_block_size num_blocks = config.num_kvcache_blocks assert num_blocks > 0 @@ -98,18 +110,28 @@ def __init__(self, config: Config): self.state_checkpoint_interval_tokens = max( 0, int(getattr(config, "state_checkpoint_interval_tokens", 0) or 0) ) - # The rolling state class: per-request groups plus a content index over - # the free ones. A checkpoint IS a free group whose content is still - # valid, so it holds no capacity of its own and never blocks admission. + checkpoint_spec = state_runtime.checkpoint_spec + self.paged_state_checkpoints: PagedStateCheckpointCoordinator | None = None + if checkpoint_spec is not None: + self.paged_state_checkpoints = PagedStateCheckpointCoordinator( + self.kv, + checkpoint_spec, + enabled=self.enable_prefix_caching + and self.num_per_req_cache_groups > 0, + ) self.state = StateGroupPool( self.num_per_req_cache_groups, - transfer=StateTransfer.from_config( - getattr(config, "state_transfer_kind", "none") or "none", - int(getattr(config, "state_fork_tokens", 0) or 0), + transfer=( + StateTransfer.none() + if self.paged_state_checkpoints is not None + else state_runtime.transfer ), hash_block_size=self.hash_block_size, enabled=self.enable_prefix_caching, ) + self._state_checkpoint_cache: StateCheckpointCache = ( + self.paged_state_checkpoints or self.state + ) # A checkpoint is filed under the content hash of the last block it # covers, so a rung that isn't a hash-block boundary can never be looked # up — the ladder would checkpoint into a void. The interval defaults to @@ -141,7 +163,7 @@ def __init__(self, config: Config): # by the state checkpoint, so it has nothing to say about hit length. # Kept plural because GDN's recurrent state is a second member the # moment it stops forking (see the state-cache protocol). - self.state_caches: tuple[StateCache, ...] = (self.state,) + self.state_caches: tuple[StateCache, ...] = (self._state_checkpoint_cache,) # The demand funnel: recorded at admission, cut for when a prefill # chunk is shortened to land on it, kept when the state pool files it. @@ -158,25 +180,23 @@ def compute_hash(cls, token_ids: list[int], prefix: int = -1): h.update(np.array(token_ids).tobytes()) return h.intdigest() - def release_state_pins(self) -> None: - """Return the previous step's resume sources to the free list. - - Called once per engine step before scheduling. A source is read by the - forward that was already issued when it is handed out again — read - directly under a fork, copied out of under a copy — and the next owner's - forward is issued after that one on the same stream, so stream ordering - covers the overlap either way. - """ + def complete_previous_state_batch(self) -> None: + """Complete state reads and copies issued by the previous batch.""" self.state.release_pins() - - def state_copies_for_batch(self) -> list[tuple[int, int]]: - """State copies the batch now being built has to issue before its - forward — checkpoints being kept and checkpoints being resumed from. - - Must be called with the batch already decided; see - `StateGroupPool.take_copies`. Always empty for a forking backend. - """ - return self.state.take_copies() + if self.paged_state_checkpoints is not None: + self.paged_state_checkpoints.complete_previous_batch() + + def take_state_maintenance_ops(self) -> StateMaintenanceOps: + """Drain state maintenance for the batch being built.""" + relocations = self.state.take_relocations() + stores = restores = () + if self.paged_state_checkpoints is not None: + stores, restores = self.paged_state_checkpoints.take_checkpoint_ops() + return StateMaintenanceOps( + relocations=relocations, + checkpoint_stores=stores, + checkpoint_restores=restores, + ) def _record_evicted(self, h: int) -> None: """A hash the block pool just dropped: report it, and settle the state. @@ -189,14 +209,30 @@ def _record_evicted(self, h: int) -> None: """ if self._event_log is not None: self._event_log.append(_make_block_removed([h])) - self.state.unindex(h) + self._state_checkpoint_cache.unindex(h) def _fresh_block(self) -> int: """Take a block for content this step is about to compute.""" + if not self._ensure_page_units(1): + raise AssertionError("No PAGE unit available for a fresh KV block") block_id = self.kv.pop() self.kv.allocate(block_id) return block_id + def _has_page_units( + self, count: int, protected_checkpoint_hash: int | None = None + ) -> bool: + if self.paged_state_checkpoints is None: + return self.kv.has_free(count) + return self.paged_state_checkpoints.has_available_units( + count, protected_hash=protected_checkpoint_hash + ) + + def _ensure_page_units(self, count: int) -> bool: + if self.paged_state_checkpoints is None: + return self.kv.has_free(count) + return self.paged_state_checkpoints.ensure_free_units(count) + def _dcp_num_blocks(self, seq_len: int) -> int: if self.dcp_world_size <= 1: return (seq_len + self.block_size - 1) // self.block_size @@ -288,13 +324,11 @@ def can_allocate(self, seq: Sequence) -> int: Caller (scheduler) passes the returned hit count to `allocate()`, avoiding a second hash pass. """ - # State cache (mamba / V4 compressor ring) has its own pre-allocated - # tensor; admission only needs a free slot index, not extra paged - # blocks. See `allocate()` for the budget reasoning. + # Active Slots are preallocated; PAGE checkpoints share the KV pool. 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._dcp_num_blocks(len(seq))): + if not self._has_page_units(self._dcp_num_blocks(len(seq))): return -1 return 0 # Step 1: compressed prefix (CSA/HCA/indexer share the block hash and @@ -341,7 +375,10 @@ def can_allocate(self, seq: Sequence) -> int: for i in range(num_cached_blocks): if self.kv.is_used(self.kv.lookup(block_hashes[i])): num_new_blocks -= 1 - if not self.kv.has_free(num_new_blocks): + protected_hash = ( + block_hashes[num_cached_blocks - 1] if num_cached_blocks else None + ) + if not self._has_page_units(num_new_blocks, protected_hash): return -1 return num_cached_blocks @@ -363,6 +400,9 @@ def allocate(self, seq: Sequence, num_cached_blocks: int = 0): block_id = self.kv.lookup(h) self.kv.claim(block_id) seq.block_table.append(block_id) + # Pin the restore before fresh blocks can evict its checkpoint. + if seq.has_per_req_cache and self.paged_state_checkpoints is not None: + self._attach_state_group(seq, h if num_cached_blocks > 0 else -1) for _ in range(num_cached_blocks, self._dcp_num_blocks(len(seq))): seq.block_table.append(self._fresh_block()) seq.num_cached_tokens = num_cached_blocks * self._hash_block_size() @@ -374,7 +414,7 @@ def allocate(self, seq: Sequence, num_cached_blocks: int = 0): # paged-block cost. The slot cap # (the state pool's free list, size = `max_num_seqs`) is the sole # admission bound for state cache. - if seq.has_per_req_cache: + if seq.has_per_req_cache and self.paged_state_checkpoints is None: self._attach_state_group(seq, h if num_cached_blocks > 0 else -1) def _attach_state_group(self, seq: Sequence, hit_hash: int) -> None: @@ -384,23 +424,30 @@ def _attach_state_group(self, seq: Sequence, hit_hash: int) -> None: start). `can_allocate` already shrank the hit to a boundary that carries a checkpoint, so a lookup miss here just means the pool is off. - Resuming shares: the checkpoint stays indexed and the request gets a - group of its own, so a second request hitting the same prefix still - finds it. How the state reaches that group is the backend's - `StateTransfer` — a fork reads the checkpoint for one forward, a copy is - handed the bytes — and the two differ by one line here. When no second - group is free the request adopts the checkpoint instead: still correct, - the state is exactly the one it wanted, it just spends the checkpoint - rather than sharing it, and under either mechanism it needs nothing. + PAGE checkpoints gather into a fresh slot; only fork checkpoints can + be adopted as request slots. A checkpoint is read-only, so several requests in one step may resume off the same one. The first takes it off the free list and the pin - covers every reader until `release_state_pins`; a later one in that same - step finds it already pinned and only needs a group to write into. + covers every reader until the previous state batch completes; a later + one in that step finds it pinned and only needs a group to write into. Adopting is then off the table — the pin means someone else's forward still has to read it, or copy out of it. """ - src = self.state.lookup(hit_hash) if hit_hash != -1 else -1 + if self.paged_state_checkpoints is not None: + dst = self.state.pop() + if hit_hash != -1 and not self.paged_state_checkpoints.begin_restore( + hit_hash, dst + ): + self.state.release(dst) + raise RuntimeError( + "gated PAGE checkpoint disappeared before state attach" + ) + seq.per_req_cache_group = dst + seq.state_fork_src = -1 + return + + src = self.state.lookup_group(hit_hash) if hit_hash != -1 else -1 if src < 0: seq.per_req_cache_group = self.state.pop() seq.state_fork_src = -1 @@ -411,10 +458,7 @@ def _attach_state_group(self, seq: Sequence, hit_hash: int) -> None: if self.state.has_free(): dst = self.state.pop() seq.per_req_cache_group = dst - if self.state.transfer.copies: - self.state.record_copy(src, dst) - else: - seq.state_fork_src = src + seq.state_fork_src = src # Held off the free list until the forward that reads it is issued. self.state.pin(src) return @@ -689,7 +733,7 @@ def checkpoint_funnel(self) -> dict[str, int]: return { "demands_recorded": self.demands_recorded, "chunks_cut_for_demand": self.chunks_cut_for_demand, - } | self.state.checkpoint_fates() + } | self._state_checkpoint_cache.checkpoint_fates() def checkpointers_at( self, @@ -905,15 +949,11 @@ def deallocate(self, seq: Sequence): seq.last_checkpoint_pos = 0 # An uncommitted checkpoint describes state in a group that is about to # go back on the free list, so the intent dies with it. - self.state.forget_pending(seq) + if self.paged_state_checkpoints is not None: + self.paged_state_checkpoints.forget_pending(seq) seq.block_table.clear() if seq.has_per_req_cache and seq.per_req_cache_group >= 0: - # Only the group the seq was writing. A checkpoint it took is - # already back on the free list under the state index; the source - # it was going to fork off is dropped here rather than left to - # `release_state_pins`, because the forward that owed the read is - # not going to happen and the group should not sit out a pass for - # a reader that no longer exists. + # No next forward will read a pending fork source after deallocation. self.state.release(seq.per_req_cache_group) self.state.drop_reader(seq.state_fork_src) seq.per_req_cache_group = -1 @@ -925,7 +965,7 @@ def can_append(self, seq: Sequence, num_new_tokens: int = 1) -> bool: ebs = self._effective_block_size() needed_blocks = (seq_len + num_new_tokens + ebs - 1) // ebs new_blocks_needed = max(0, needed_blocks - current_blocks) - return self.kv.has_free(new_blocks_needed) + return self._has_page_units(new_blocks_needed) def may_append(self, seq: Sequence, num_new_tokens: int = 1): # Note: in disaggregated (P/D) mode the scheduler skips this call on @@ -960,6 +1000,7 @@ def clear_cache(self) -> None: they remain valid via their block_table refs, just unhashable for future requests.""" self.kv.clear_index() + self._state_checkpoint_cache.clear_index() if self._event_log is not None: self._event_log.append(_make_all_cleared()) diff --git a/atom/model_engine/block_pool.py b/atom/model_engine/block_pool.py index eac173dbc1..326b68ed1b 100644 --- a/atom/model_engine/block_pool.py +++ b/atom/model_engine/block_pool.py @@ -2,7 +2,7 @@ # Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. from collections import OrderedDict -from collections.abc import Callable +from collections.abc import Callable, Hashable, Iterable from dataclasses import dataclass from heapq import heapify, heappop, heappush @@ -84,6 +84,8 @@ def __init__( self._cached: OrderedDict[int, None] = OrderedDict() self._free: set[int] = set(range(num_blocks)) self._used: set[int] = set() + # Raw PAGE units reserved by multi-unit objects such as state checkpoints. + self._raw_unit_owner: dict[int, tuple[Hashable, int]] = {} # ------------------------------- counts -------------------------------- # @property @@ -215,6 +217,10 @@ def claim(self, block_id: int) -> Block: return block def free(self, block_id: int) -> None: + if block_id in self._raw_unit_owner: + raise AssertionError( + f"block {block_id} is a reserved raw unit; use release_units" + ) block = self.blocks[block_id] block.ref_count -= 1 if block.ref_count: @@ -232,6 +238,39 @@ def free(self, block_id: int) -> None: self._vacant = [b for b in self._free if self.blocks[b].hash == -1] heapify(self._vacant) + def reserve_units(self, count: int, owner: Hashable) -> list[int] | None: + """Reserve arbitrary PAGE-sized units for raw storage.""" + if count < 0: + raise ValueError(f"unit count must be non-negative, got {count}") + if owner is None: + raise ValueError("a raw-unit reservation needs an owner") + if not self.has_free(count): + return None + unit_ids: list[int] = [] + for piece_index in range(count): + block_id = self.pop() + self.allocate(block_id) + self._raw_unit_owner[block_id] = (owner, piece_index) + unit_ids.append(block_id) + return unit_ids + + def release_units(self, unit_ids: Iterable[int], owner: Hashable) -> None: + """Release a complete raw-unit reservation back to the PAGE pool.""" + ids = list(unit_ids) + if len(ids) != len(set(ids)): + raise ValueError("a raw-unit release contains duplicate ids") + # Validate ownership before releasing any unit. + for piece_index, block_id in enumerate(ids): + actual = self._raw_unit_owner.get(block_id) + expected = (owner, piece_index) + if actual != expected: + raise AssertionError( + f"raw unit {block_id} belongs to {actual!r}, not {expected!r}" + ) + for block_id in ids: + del self._raw_unit_owner[block_id] + self.free(block_id) + # ------------------------------ resizing ------------------------------- # def extend(self, count: int) -> int: """Grow the pool by up to `count` blocks; returns how many it took. @@ -262,6 +301,9 @@ def retire_top(self) -> BlockRetirement | None: top = self.num_blocks - 1 if top < 0: return None + # Raw units cannot move without updating their owning record. + if top in self._raw_unit_owner: + return None if top in self._free: self._take_named(top) self._unindex(top) diff --git a/atom/model_engine/engine_core.py b/atom/model_engine/engine_core.py index c1522fe2db..fad4138ffc 100644 --- a/atom/model_engine/engine_core.py +++ b/atom/model_engine/engine_core.py @@ -18,6 +18,7 @@ from atom.model_engine.engine_utility import EngineUtilityHandler from atom.model_engine.scheduler import DecodeScheduler, PrefillScheduler, Scheduler from atom.model_engine.sequence import Sequence, SequenceStatus, get_exit_sequence +from atom.model_engine.state_runtime import StateRuntime from atom.utils import ( envs, init_exit_handler, @@ -93,8 +94,7 @@ def __init__(self, config: Config, input_address: str, output_address: str): # adding an architecture never touches this line. config.pool_entries = block_info.get("pool_entries", {}) config.pool_entries_per_req = block_info.get("pool_entries_per_req", {}) - config.state_transfer_kind = block_info.get("state_transfer_kind", "none") - config.state_fork_tokens = block_info.get("state_fork_tokens", 0) + self.state_runtime = StateRuntime.from_wire(block_info["state_runtime"]) ret = self.runner_mgr.call_func( "allocate_kv_cache", num_blocks, wait_out=True ) @@ -123,7 +123,10 @@ def __init__(self, config: Config, input_address: str, output_address: str): # consumers can reference it before DecodeEngineCore creates the real one. self.scheduler = None if not config.disagg_is_decode: - self.scheduler = Scheduler(config) + self.scheduler = Scheduler( + config, + state_runtime=self.state_runtime, + ) self.kv_transfer_enabled = bool(config.kv_transfer_config) if self.kv_transfer_enabled: @@ -917,7 +920,9 @@ def __init__(self, config: Config, input_address: str, output_address: str): # --- Create DecodeScheduler now that num_kvcache_blocks is set --- self.scheduler = DecodeScheduler( - config, disagg_cu_shm_name=config.disagg_cu_shm_name + config, + disagg_cu_shm_name=config.disagg_cu_shm_name, + state_runtime=self.state_runtime, ) # EngineUtilityHandler was built in super().__init__() with scheduler=None # (decode defers scheduler creation); wire the real one in for MTP stats. diff --git a/atom/model_engine/model_runner.py b/atom/model_engine/model_runner.py index 732472afda..6e8a3efa6e 100644 --- a/atom/model_engine/model_runner.py +++ b/atom/model_engine/model_runner.py @@ -38,9 +38,12 @@ recv_intermediate_tensors, ) from atom.kv_transfer.disaggregation import KVConnectorOutput +from atom.model_engine.kv_block import STATE_SLOT_CLASS +from atom.model_engine.page_unit_checkpoint import PagedStateCheckpointSpec from atom.model_engine.run_labels import build_run_label from atom.model_engine.scheduler import ScheduledBatch, ScheduledBatchOutput from atom.model_engine.sequence import Sequence, SequenceStatus, SequenceType +from atom.model_engine.state_runtime import StateRuntime from atom.model_loader.loader import load_model from atom.model_ops.attentions.sub_pool_spec import ( InsufficientPoolBudget, @@ -740,6 +743,7 @@ def __init__(self, rank: int, config: Config): # builder through paths that ask for their entry counts, and those must # read 0 ("no pool yet") rather than trip over a missing attribute. self.pool_plan = PoolPlan.empty() + self.state_runtime = StateRuntime() # Sanity-check: any builder that allocates a per-request cache must # have its model_type listed in `InputOutputProcessor`'s # `per_req_cache_model_types` set; otherwise sequences will be @@ -1564,7 +1568,7 @@ def _estimate_cudagraph_overhead(self): ) return int(overhead) - def get_num_blocks(self) -> dict[str, int]: + def get_num_blocks(self) -> dict[str, object]: torch.set_default_device(self.device) config = self.config hf_config = config.hf_config @@ -1665,11 +1669,44 @@ def get_num_blocks(self) -> dict[str, int]: self.pool_plan = plan config.pool_entries = dict(plan.entries) config.pool_entries_per_req = dict(plan.entries_per_req) - # Two scalars rather than the StateTransfer itself: this travels to the - # engine process in a plain dict, where `StateGroupPool` rebuilds it. + # Keep runtime state metadata out of Config. transfer = self.attn_metadata_builder.state_transfer() - config.state_transfer_kind = transfer.kind - config.state_fork_tokens = transfer.fork_tokens + uses_paged_state = transfer.copies + if uses_paged_state and config.pipeline_parallel_size > 1: + raise RuntimeError( + "PAGE-backed state checkpoints do not yet support pipeline " + "parallelism: every stage must first agree on one atomic " + "checkpoint/unit ownership transaction" + ) + if uses_paged_state and config.enable_rapidserve: + raise RuntimeError( + "PAGE-backed state checkpoints do not yet support RapidServe " + "prefill/decode disaggregation" + ) + checkpoint_spec = None + if uses_paged_state: + if plan.paged_class is None: + raise RuntimeError( + "PAGE-backed state checkpoints require a PAGE sub-pool" + ) + checkpoint_spec = PagedStateCheckpointSpec( + page_unit_bytes=int(plan.entry_bytes[plan.paged_class]), + slot_bytes=int(plan.entry_bytes[STATE_SLOT_CLASS]), + layout_id=transfer.paged_layout_id, + ) + logger.info( + "PAGE-backed state checkpoints enabled: unit_bytes=%d, " + "slot_bytes=%d, units_per_checkpoint=%d, layout=%s", + checkpoint_spec.page_unit_bytes, + checkpoint_spec.slot_bytes, + checkpoint_spec.units_per_checkpoint, + checkpoint_spec.layout_id, + ) + state_runtime = StateRuntime( + transfer=transfer, + checkpoint_spec=checkpoint_spec, + ) + self.state_runtime = state_runtime for name in sorted(plan.entries): logger.info( f"sub-pool {name}: entries={plan.entries[name]}, " @@ -1693,11 +1730,7 @@ def get_num_blocks(self) -> dict[str, int]: # Concurrent-capacity table: at each context-length percentage of # max_model_len, how many requests can simultaneously hold their # KV in the pool. Per-req block usage = ceil(ctx_len/block_size). - # STATE classes sit in their own reservation (already excluded from - # the paged count at sizing time), so they add no per-block cost and - # never bind either: sizing reserves every STATE floor at exactly - # `max_num_seqs` requests' worth, so the request cap is max_num_seqs - # and the paged pool is the only thing that can run out first. + # Active Slots are reserved; PAGE checkpoints borrow from the paged pool. max_model_len = config.max_model_len cap = config.max_num_seqs pct_lines = [] @@ -1742,8 +1775,7 @@ def get_num_blocks(self) -> dict[str, int]: "num_kvcache_blocks": num_kvcache_blocks, "pool_entries": dict(plan.entries), "pool_entries_per_req": dict(plan.entries_per_req), - "state_transfer_kind": config.state_transfer_kind, - "state_fork_tokens": config.state_fork_tokens, + "state_runtime": state_runtime.to_wire(), } def allocate_kv_cache(self, num_kvcache_blocks): @@ -4102,11 +4134,20 @@ def _kv_budget_extra_reserve(self, total_bytes: int) -> int: safety_margin = int(total_bytes * 0.02) return 4 * safety_margin - def get_num_blocks(self) -> dict[str, int]: + def get_num_blocks(self) -> dict[str, object]: # Decode in disagg mode owns no GPU memory — kvcache is imported from # prefill. if self.config.disagg_is_decode: - return {"num_kvcache_blocks": 0} + transfer = self.attn_metadata_builder.state_transfer() + if transfer.copies: + raise RuntimeError( + "PAGE-backed state checkpoints do not yet support RapidServe " + "prefill/decode disaggregation" + ) + return { + "num_kvcache_blocks": 0, + "state_runtime": StateRuntime(transfer=transfer).to_wire(), + } return super().get_num_blocks() def allocate_kv_cache(self, num_kvcache_blocks): diff --git a/atom/model_engine/page_unit_checkpoint.py b/atom/model_engine/page_unit_checkpoint.py new file mode 100644 index 0000000000..575ddb65e6 --- /dev/null +++ b/atom/model_engine/page_unit_checkpoint.py @@ -0,0 +1,404 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""State checkpoints backed by arbitrary PAGE-sized physical units.""" + +from __future__ import annotations + +from collections import OrderedDict +from collections.abc import Mapping +from dataclasses import dataclass + +from atom.model_engine.block_pool import BlockPool +from atom.model_engine.sequence import Sequence + +COPYING = "COPYING" +READY = "READY" +EVICTING = "EVICTING" + + +@dataclass(frozen=True) +class PagedStateCheckpointSpec: + """Runtime geometry for PAGE-backed state checkpoints.""" + + page_unit_bytes: int + slot_bytes: int + layout_id: str + + def __post_init__(self) -> None: + for name, value in ( + ("page_unit_bytes", self.page_unit_bytes), + ("slot_bytes", self.slot_bytes), + ): + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer") + if not isinstance(self.layout_id, str) or not self.layout_id: + raise ValueError("paged state checkpoints need a non-empty layout id") + + @property + def units_per_checkpoint(self) -> int: + return (self.slot_bytes + self.page_unit_bytes - 1) // self.page_unit_bytes + + def to_wire(self) -> dict[str, int | str]: + return { + "page_unit_bytes": self.page_unit_bytes, + "slot_bytes": self.slot_bytes, + "layout_id": self.layout_id, + } + + @classmethod + def from_wire(cls, wire: object) -> PagedStateCheckpointSpec: + if not isinstance(wire, Mapping): + raise TypeError("paged state checkpoint spec must be a mapping") + expected = {"page_unit_bytes", "slot_bytes", "layout_id"} + if set(wire) != expected: + raise ValueError( + "invalid paged state checkpoint spec fields: " + f"expected={sorted(expected)}, got={sorted(wire)}" + ) + return cls( + page_unit_bytes=wire["page_unit_bytes"], # type: ignore[arg-type] + slot_bytes=wire["slot_bytes"], # type: ignore[arg-type] + layout_id=wire["layout_id"], # type: ignore[arg-type] + ) + + +@dataclass(frozen=True) +class CheckpointStoreOp: + """Scatter one contiguous Active Slot into ordered PAGE units.""" + + src_slot: int + unit_ids: tuple[int, ...] + total_bytes: int + layout_id: str + + +@dataclass(frozen=True) +class CheckpointRestoreOp: + """Gather one ordered PAGE-unit image into an Active Slot.""" + + dst_slot: int + unit_ids: tuple[int, ...] + total_bytes: int + layout_id: str + + +@dataclass +class CheckpointRecord: + prefix_hash: int + unit_ids: tuple[int, ...] + state: str = COPYING + pin_count: int = 0 + + +class PageUnitCheckpointStore: + """Content index and ownership table for split state images.""" + + def __init__( + self, + pool: BlockPool, + spec: PagedStateCheckpointSpec, + ): + self.pool = pool + self.spec = spec + + self.hash_to_checkpoint: dict[int, int] = {} + self.records: dict[int, CheckpointRecord] = {} + self._pending_by_hash: dict[int, int] = {} + self._lru: OrderedDict[int, None] = OrderedDict() + self._inflight_stores: list[int] = [] + self._queued_restores: list[tuple[int, CheckpointRestoreOp]] = [] + self._inflight_restores: list[int] = [] + self._next_checkpoint_id = 0 + self.evictions = 0 + + @property + def units_per_checkpoint(self) -> int: + return self.spec.units_per_checkpoint + + def lookup(self, prefix_hash: int) -> int: + checkpoint_id = self.hash_to_checkpoint.get(prefix_hash, -1) + record = self.records.get(checkpoint_id) + if record is None or record.state != READY: + return -1 + return checkpoint_id + + def contains(self, prefix_hash: int) -> bool: + return self.lookup(prefix_hash) >= 0 + + def contains_or_pending(self, prefix_hash: int) -> bool: + return self.contains(prefix_hash) or prefix_hash in self._pending_by_hash + + def _new_identity(self) -> int: + checkpoint_id = self._next_checkpoint_id + self._next_checkpoint_id += 1 + return checkpoint_id + + def has_available_units( + self, count: int, protected_hash: int | None = None + ) -> bool: + if count <= self.pool.num_free: + return True + protected = self.lookup(protected_hash) if protected_hash is not None else -1 + reclaimable = sum( + len(self.records[cid].unit_ids) + for cid in self._lru + if cid != protected + if self.records[cid].state == READY and self.records[cid].pin_count == 0 + ) + return self.pool.num_free + reclaimable >= count + + def ensure_free_units(self, count: int) -> bool: + while self.pool.num_free < count: + victim = next( + ( + cid + for cid in self._lru + if self.records[cid].state == READY + and self.records[cid].pin_count == 0 + ), + -1, + ) + if victim < 0: + return False + self._evict(victim) + return True + + def begin_store(self, prefix_hash: int, src_slot: int) -> CheckpointStoreOp | None: + if self.lookup(prefix_hash) >= 0 or prefix_hash in self._pending_by_hash: + return None + needed = self.units_per_checkpoint + if not self.ensure_free_units(needed): + return None + + checkpoint_id = self._new_identity() + owner = ("state-checkpoint", checkpoint_id) + unit_ids = self.pool.reserve_units(needed, owner) + if unit_ids is None: + return None + record = CheckpointRecord( + prefix_hash=prefix_hash, + unit_ids=tuple(unit_ids), + ) + self.records[checkpoint_id] = record + self._pending_by_hash[prefix_hash] = checkpoint_id + self._inflight_stores.append(checkpoint_id) + return CheckpointStoreOp( + src_slot=src_slot, + unit_ids=record.unit_ids, + total_bytes=self.spec.slot_bytes, + layout_id=self.spec.layout_id, + ) + + def begin_restore( + self, prefix_hash: int, dst_slot: int + ) -> CheckpointRestoreOp | None: + checkpoint_id = self.lookup(prefix_hash) + if checkpoint_id < 0: + return None + record = self.records[checkpoint_id] + record.pin_count += 1 + self._lru.move_to_end(checkpoint_id) + op = CheckpointRestoreOp( + dst_slot=dst_slot, + unit_ids=record.unit_ids, + total_bytes=self.spec.slot_bytes, + layout_id=self.spec.layout_id, + ) + self._queued_restores.append((checkpoint_id, op)) + return op + + def take_restore_ops(self) -> tuple[CheckpointRestoreOp, ...]: + queued, self._queued_restores = self._queued_restores, [] + self._inflight_restores.extend(checkpoint_id for checkpoint_id, _ in queued) + return tuple(op for _, op in queued) + + def cancel_queued_restore(self, dst_slot: int) -> None: + kept: list[tuple[int, CheckpointRestoreOp]] = [] + for checkpoint_id, op in self._queued_restores: + if op.dst_slot == dst_slot: + self._release_restore_pin(checkpoint_id) + else: + kept.append((checkpoint_id, op)) + self._queued_restores = kept + + def complete_inflight(self) -> None: + stores, self._inflight_stores = self._inflight_stores, [] + for checkpoint_id in stores: + record = self.records.get(checkpoint_id) + if record is None: + continue + if self._pending_by_hash.get(record.prefix_hash) == checkpoint_id: + del self._pending_by_hash[record.prefix_hash] + if record.state == EVICTING: + self._release_record(checkpoint_id) + continue + if record.state != COPYING: + continue + # Publish only after the scatter has ridden a batch. + if self.lookup(record.prefix_hash) >= 0: + self._release_record(checkpoint_id) + continue + record.state = READY + self.hash_to_checkpoint[record.prefix_hash] = checkpoint_id + self._lru[checkpoint_id] = None + + restores, self._inflight_restores = self._inflight_restores, [] + for checkpoint_id in restores: + self._release_restore_pin(checkpoint_id) + + def _release_restore_pin(self, checkpoint_id: int) -> None: + record = self.records.get(checkpoint_id) + if record is None: + return + if record.pin_count <= 0: + raise AssertionError("checkpoint restore pin underflow") + record.pin_count -= 1 + if record.state == EVICTING and record.pin_count == 0: + self._release_record(checkpoint_id) + + def unindex(self, prefix_hash: int) -> bool: + checkpoint_id = self.hash_to_checkpoint.pop(prefix_hash, -1) + if checkpoint_id < 0: + checkpoint_id = self._pending_by_hash.pop(prefix_hash, -1) + if checkpoint_id < 0: + return False + record = self.records.get(checkpoint_id) + if record is None: + return False + record.state = EVICTING + self._lru.pop(checkpoint_id, None) + # Keep units alive while a queued GPU writer can still access them. + if checkpoint_id not in self._inflight_stores and record.pin_count == 0: + self._release_record(checkpoint_id) + return True + + def clear(self) -> None: + self.hash_to_checkpoint.clear() + self._pending_by_hash.clear() + self._lru.clear() + inflight_stores = set(self._inflight_stores) + for checkpoint_id in list(self.records): + record = self.records[checkpoint_id] + record.state = EVICTING + if checkpoint_id not in inflight_stores and record.pin_count == 0: + self._release_record(checkpoint_id) + + def _evict(self, checkpoint_id: int) -> None: + record = self.records[checkpoint_id] + if record.state != READY or record.pin_count: + raise AssertionError("only an unpinned READY checkpoint is evictable") + if self.hash_to_checkpoint.get(record.prefix_hash) == checkpoint_id: + del self.hash_to_checkpoint[record.prefix_hash] + record.state = EVICTING + self._lru.pop(checkpoint_id, None) + self._release_record(checkpoint_id) + self.evictions += 1 + + def _release_record(self, checkpoint_id: int) -> None: + record = self.records.pop(checkpoint_id) + self._lru.pop(checkpoint_id, None) + if self.hash_to_checkpoint.get(record.prefix_hash) == checkpoint_id: + del self.hash_to_checkpoint[record.prefix_hash] + if self._pending_by_hash.get(record.prefix_hash) == checkpoint_id: + del self._pending_by_hash[record.prefix_hash] + self.pool.release_units(record.unit_ids, ("state-checkpoint", checkpoint_id)) + + +class PagedStateCheckpointCoordinator: + """Schedules PAGE-backed checkpoints for per-request state.""" + + successor_room = 0.0 + + def __init__( + self, + pool: BlockPool, + spec: PagedStateCheckpointSpec, + enabled: bool, + ) -> None: + self.enabled = enabled + self.store = PageUnitCheckpointStore(pool, spec) + self._pending: dict[int, tuple[Sequence, int]] = {} + self._store_ops: list[CheckpointStoreOp] = [] + self.checkpoints_kept = 0 + self.checkpoints_dropped = 0 + self.checkpoints_orphaned = 0 + + def applies(self, seq: Sequence) -> bool: + return self.enabled and seq.has_per_req_cache + + def resumable_hit( + self, + seq: Sequence, + hit: int, + block_hashes: list[int], + assume_checkpointed: bool = False, + ) -> int: + if not self.applies(seq): + return hit + for i in range(hit - 1, -1, -1): + if assume_checkpointed or self.store.contains(block_hashes[i]): + return i + 1 + return 0 + + def checkpoint(self, seq: Sequence, boundary_blocks: int, h: int) -> None: + del boundary_blocks + if self.applies(seq) and seq.per_req_cache_group >= 0: + self._pending[id(seq)] = (seq, h) + + def forget_pending(self, seq: Sequence) -> None: + self._pending.pop(id(seq), None) + self.store.cancel_queued_restore(seq.per_req_cache_group) + + def begin_restore(self, h: int, dst_slot: int) -> bool: + return self.store.begin_restore(h, dst_slot) is not None + + def take_checkpoint_ops( + self, + ) -> tuple[tuple[CheckpointStoreOp, ...], tuple[CheckpointRestoreOp, ...]]: + pending, self._pending = self._pending, {} + for seq, h in pending.values(): + src_slot = seq.per_req_cache_group + if src_slot < 0 or self.store.contains_or_pending(h): + continue + op = self.store.begin_store(h, src_slot) + if op is None: + self.checkpoints_dropped += 1 + continue + self._store_ops.append(op) + self.checkpoints_kept += 1 + stores, self._store_ops = self._store_ops, [] + return tuple(stores), self.store.take_restore_ops() + + def complete_previous_batch(self) -> None: + self.store.complete_inflight() + + def has_available_units( + self, count: int, protected_hash: int | None = None + ) -> bool: + return self.store.has_available_units(count, protected_hash) + + def ensure_free_units(self, count: int) -> bool: + return self.store.ensure_free_units(count) + + def unindex(self, h: int) -> None: + pending_ids = [ + seq_id for seq_id, (_, pending_h) in self._pending.items() if pending_h == h + ] + for seq_id in pending_ids: + del self._pending[seq_id] + removed = self.store.unindex(h) + if pending_ids or removed: + self.checkpoints_orphaned += 1 + + def clear_index(self) -> None: + self._pending.clear() + self.store.clear() + + def checkpoint_fates(self) -> dict[str, int]: + return { + "checkpoints_kept": self.checkpoints_kept, + "checkpoints_dropped": self.checkpoints_dropped, + "checkpoints_evicted": self.store.evictions, + "checkpoints_orphaned": self.checkpoints_orphaned, + } diff --git a/atom/model_engine/scheduler.py b/atom/model_engine/scheduler.py index de5c1629c3..b0a81b4304 100644 --- a/atom/model_engine/scheduler.py +++ b/atom/model_engine/scheduler.py @@ -30,6 +30,11 @@ from atom.model_engine.block_manager import BlockManager from atom.model_engine.request import RequestOutput from atom.model_engine.sequence import Sequence, SequenceStatus, SequenceType +from atom.model_engine.state_runtime import ( + DEFAULT_STATE_RUNTIME, + StateMaintenanceOps, + StateRuntime, +) from atom.utils import envs logger = logging.getLogger("atom") @@ -318,9 +323,7 @@ class ScheduledBatch: num_spec_step: Number of speculative decode steps (0 = disabled). scheduled_spec_decode_tokens: Draft token IDs per request for speculative decoding (must not use a mutable default). - state_copy_pairs: (src, dst) per-request state groups this batch's - forward must duplicate before running (`BlockManager - .state_copies_for_batch`). + state_maintenance_ops: State moves that must execute before this batch. """ def __init__( @@ -343,7 +346,7 @@ def __init__( num_cached_tokens: list[int] | None = None, is_final_chunk: list[bool] | None = None, next_token_ids: list[int] | None = None, - state_copy_pairs: list[tuple[int, int]] | None = None, + state_maintenance_ops: StateMaintenanceOps | None = None, ): if scheduled_spec_decode_tokens is None: scheduled_spec_decode_tokens = {} @@ -379,13 +382,12 @@ def __init__( for seq in seqs.values() if seq.has_per_req_cache and seq.per_req_cache_group >= 0 ] - # (src, dst) state groups this batch's forward must duplicate before it - # runs — the copy twin of `state_fork_srcs`, for backends that checkpoint - # by copying. Not per-seq: a copy is between two pool slots and needs no - # alignment with anything else on the batch. Passed in rather than read - # off the seqs because both halves (checkpoint taken, checkpoint resumed - # from) accumulate in the pool during this pass. - self.state_copy_pairs = state_copy_pairs or [] + # Physical moves are drained once per real batch. + self.state_maintenance_ops = ( + state_maintenance_ops + if state_maintenance_ops is not None + else StateMaintenanceOps() + ) self.top_ks = np.asarray([seq.top_k for seq in seqs.values()], dtype=np.int32) self.top_ps = np.asarray([seq.top_p for seq in seqs.values()], dtype=np.float32) # True if any seq in the batch is a fan-out child (SamplingParams.n>1) @@ -610,7 +612,12 @@ class Scheduler: :meth:`_update_from_kv_xfer_finished` (both sides). """ - def __init__(self, config: Config): + def __init__( + self, + config: Config, + *, + state_runtime: StateRuntime = DEFAULT_STATE_RUNTIME, + ): self.max_num_seqs = config.max_num_seqs self.max_num_batched_tokens = config.max_num_batched_tokens self.long_prefill_token_threshold = config.long_prefill_token_threshold @@ -618,7 +625,10 @@ def __init__(self, config: Config): self.bos_token_id = config.bos_token_id self.eos_token_id = config.eos_token_id self.stop_token_ids = config.stop_token_ids - self.block_manager = BlockManager(config) + self.block_manager = BlockManager( + config, + state_runtime=state_runtime, + ) self.waiting: deque[Sequence] = deque() self.running: deque[Sequence] = deque() self.config = config @@ -1044,8 +1054,8 @@ def schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]]: """ self._schedule_tick += 1 # Sources borrowed by the previous batch: its forward has been issued, - # so they can go back on the free list (see release_state_pins). - self.block_manager.release_state_pins() + # so they can go back on the free list. + self.block_manager.complete_previous_state_batch() scheduled_seqs = {} num_seqs_prefill = 0 num_batched_tokens = 0 @@ -1366,7 +1376,7 @@ def schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]]: num_cached_tokens=num_cached_tokens_list, is_final_chunk=is_final_chunk, next_token_ids=next_token_ids, - state_copy_pairs=self.block_manager.state_copies_for_batch(), + state_maintenance_ops=self.block_manager.take_state_maintenance_ops(), ) self._consume_state_forks(scheduled_seqs) @@ -1486,14 +1496,11 @@ def schedule(self) -> tuple[ScheduledBatch, dict[int, Sequence]]: scheduled_spec_decode_tokens=scheduled_spec_decode_tokens, remote_kv_block_ids=sorted(remote_kv_blocks) if remote_kv_blocks else [], remote_kv_seq_blocks=remote_kv_seq_blocks, - # An empty batch is not forwarded (`engine_core` skips on zero - # req_ids), so draining here would file the destinations in the - # index and then never issue the copies that fill them — a resumer - # would read the previous occupant's state, the exact #1417 shape - # the copies exist to prevent. Leave them pending for the next - # batch that actually runs. - state_copy_pairs=( - self.block_manager.state_copies_for_batch() if scheduled_seqs else () + # An empty batch cannot execute queued maintenance. + state_maintenance_ops=( + self.block_manager.take_state_maintenance_ops() + if scheduled_seqs + else None ), ) self._consume_state_forks(scheduled_seqs) @@ -2811,8 +2818,17 @@ class DecodeScheduler(Scheduler): running. schedule() only schedules the running queue as decode batches. """ - def __init__(self, config: Config, disagg_cu_shm_name: str = ""): - super().__init__(config) + def __init__( + self, + config: Config, + disagg_cu_shm_name: str = "", + *, + state_runtime: StateRuntime = DEFAULT_STATE_RUNTIME, + ): + super().__init__( + config, + state_runtime=state_runtime, + ) # seq_id → Sequence; blocks allocated, BlockAssignment sent, awaiting PrefillDone. self.prefill_waiting: dict[int, Sequence] = {} self.prefill_done: deque[Sequence] = deque() @@ -2902,7 +2918,7 @@ def schedule(self): # through the same `block_manager.allocate` and the same `postprocess`, # so it owes the state pool the same two hooks. Without this one the # pins taken by every resume accumulate forever and admission starves. - self.block_manager.release_state_pins() + self.block_manager.complete_previous_state_batch() prefill_finished = False while self.prefill_done: @@ -2983,10 +2999,7 @@ def schedule(self): num_spec_step=self.mtp_k, scheduled_spec_decode_tokens=scheduled_spec_decode_tokens, cu_stream_fraction=self.cu_fraction, - # The other half of the pair above: queued copies have to reach - # a batch or the group they were filed under holds the previous - # occupant's state. - state_copy_pairs=self.block_manager.state_copies_for_batch(), + state_maintenance_ops=self.block_manager.take_state_maintenance_ops(), ), scheduled_seqs, ) diff --git a/atom/model_engine/sequence.py b/atom/model_engine/sequence.py index cc96a6a1e9..ca139f97e3 100644 --- a/atom/model_engine/sequence.py +++ b/atom/model_engine/sequence.py @@ -130,14 +130,6 @@ def __init__( # exactly one forward. -1 = read and write the same group, the case for # every step in between. self.state_fork_src = -1 - # Content hash of a boundary this seq's last forward landed on and whose - # state is worth keeping, for state classes that checkpoint by copying - # (`StateGroupPool.checkpoint`). The copy needs a forward to issue it, so - # the intent outlives the step that formed it; `StateGroupPool.take_copies` - # turns it into a destination group and a copy pair when the next batch - # is built. - # -1 = nothing pending, which is also what `deallocate` restores. - self.pending_checkpoint = -1 self.temperature = sampling_params.temperature self.top_k = sampling_params.top_k self.top_p = sampling_params.top_p diff --git a/atom/model_engine/state_cache.py b/atom/model_engine/state_cache.py index b009531bf0..d5f2f420a5 100644 --- a/atom/model_engine/state_cache.py +++ b/atom/model_engine/state_cache.py @@ -93,3 +93,13 @@ def checkpoint(self, seq: Sequence, boundary_blocks: int, h: int) -> None: nothing, and the hit it later declines is the only consequence. """ ... + + +class StateCheckpointCache(StateCache, Protocol): + """State cache lifecycle owned by BlockManager.""" + + def unindex(self, h: int) -> None: ... + + def clear_index(self) -> None: ... + + def checkpoint_fates(self) -> dict[str, int]: ... diff --git a/atom/model_engine/state_pool.py b/atom/model_engine/state_pool.py index c9195ecdcd..638ebfc725 100644 --- a/atom/model_engine/state_pool.py +++ b/atom/model_engine/state_pool.py @@ -4,83 +4,8 @@ from collections import deque from dataclasses import dataclass from heapq import heapify, heappop, heappush -from math import inf -# `StateTransfer.kind` values. Plain strings rather than an enum because the -# choice crosses a process boundary inside a dict (ModelRunner's `block_info`), -# where a scalar survives and a class does not. -FORK = "fork" -COPY = "copy" -NONE = "none" - - -@dataclass(frozen=True) -class StateTransfer: - """How a backend hands one request's state over to another group. - - Three answers, and every checkpoint decision downstream follows from which: - - `none()` no per-request state, or none that can be handed over at all. - Nothing is ever checkpointed and prefix hits shrink to 0. - `fork(n)` the state rolls. The owner gives its group to the index and - takes a fresh one, reading the old and writing the new for - exactly one forward — which has to leave the new group - self-contained, and that takes `n` committed tokens. - `copy()` one request's state is a contiguous byte range another group - can be handed a duplicate of. Nothing is given away, so - nothing downstream has to cooperate: no successor forward, and - the resuming side is handed a duplicate too. - - The two mechanisms are not interchangeable, and which one a backend can - offer decides where it may checkpoint. A fork's contract binds the *next* - forward, so it can only be taken where that forward is known to be long - enough — true on a prompt, false during generation, where a step commits - `1 + accepted_drafts` and acceptance is not knowable in advance. That is why - DeepSeek-V4 copies: it is the only way to checkpoint at a decode boundary. - See `/app/logs_claude/verify_v4_min_fork.py` for the arithmetic. - - These used to be one integer, `min_fork_tokens`, with 0 spelling `none()` — - which is exactly the value `copy()` has to report, so the two were - indistinguishable. Splitting the kind out is what lets a backend say "no - successor needed" without saying "no state". - """ - - kind: str - fork_tokens: int = 0 - - @classmethod - def none(cls) -> "StateTransfer": - return cls(NONE) - - @classmethod - def fork(cls, tokens: int) -> "StateTransfer": - assert tokens > 0, "a fork binds its successor forward; use none()" - return cls(FORK, tokens) - - @classmethod - def copy(cls) -> "StateTransfer": - return cls(COPY) - - @classmethod - def from_config(cls, kind: str, fork_tokens: int) -> "StateTransfer": - """Rebuild from the two scalars that crossed the process boundary.""" - if kind == FORK: - return cls.fork(fork_tokens) - assert kind in (COPY, NONE), f"unknown state transfer kind {kind!r}" - return cls(kind) - - @property - def copies(self) -> bool: - return self.kind == COPY - - @property - def forks(self) -> bool: - return self.kind == FORK - - @property - def successor_room(self) -> float: - """`StateCache.successor_room` for a class transferred this way.""" - return inf if self.kind == NONE else float(self.fork_tokens) +from atom.model_engine.state_runtime import StateTransfer @dataclass(frozen=True) @@ -101,52 +26,7 @@ class GroupRetirement: class StateGroupPool: - """Per-request state groups, plus a content index over the free ones. - - A *group* is what one request occupies in the pre-allocated state tensor: - `entries // entries_per_req` contiguous indices (GDN conv+ssm, the - DeepSeek-V4 compressor ring and sliding window). This pool owns the free - list, so it is the single answer to "can one more request be admitted". - - A per-request state cannot be rebuilt from cached KV blocks: the cache holds - the compressor's *output*, the state is its rolling *input* window. So a - prefix-cache hit is only recoverable up to a boundary where somebody - checkpointed the state — which is what the index over the free list - answers. - - Capacity model: a checkpoint is a group sitting on the free list with its - content still valid, filed under the content hash of the last block it - covers. This is the KV block pool's lazy-eviction model (`pop` drops the - hash at hand-out time, not at free time) applied to state groups. The index - therefore holds nothing back — `pop` invalidates whatever it hands out, so a - checkpoint can never shrink admission, and under full concurrency the - checkpoint set drains to empty on its own. - - *How* a group reaches the index is the backend's `StateTransfer`, and it is - the only thing that differs between the two mechanisms this class runs: - - `fork` the owner gives its group away and takes a fresh one, so the - checkpoint costs no bytes but binds the very next forward, which - has to leave the replacement self-contained (`min_fork_tokens`). - `copy` the state is a byte range, so a duplicate goes to the index and - the owner is not disturbed at all. Nothing is bound: no successor - forward, and the resuming side copies rather than forking too. - - Both meet the same index and the same free list. Under `copy` the bytes are - moved by a forward, so this class only schedules the pairs (`take_copies`) - and the next batch issues them. - - The count of groups is not fixed for life: `extend` and `retire_top` move it - when the state pool's share of the byte budget changes. Retiring is - index-forced but its cost is not — see `retire_top`. - - Vocabulary: *checkpoint* is the state sense throughout — a boundary this - class kept resumable. *Publish* is reserved for a block entering the - content-addressed KV index (`BlockManager.hash_blocks`, the KV events). - - `enabled` covers the *index* only. The free list stays live either way: - admission needs a group whether or not anything is ever checkpointed. - """ + """Own Active Slot allocation and the fork-checkpoint index.""" def __init__( self, @@ -158,12 +38,12 @@ def __init__( self.enabled: bool = enabled and num_groups > 0 self.num_groups: int = num_groups self.transfer: StateTransfer = transfer or StateTransfer.none() - # Committed tokens the forward after a fork must cover for the new group - # to come out self-contained. 0 under `copy`, where the destination is - # complete the moment the copy lands and no forward is involved. + # Committed tokens needed to make a fork destination self-contained. self.min_fork_tokens: int = self.transfer.fork_tokens self.successor_room: float = self.transfer.successor_room self.hash_block_size: int = hash_block_size + if self.transfer.copies: + raise ValueError("PAGE-copy checkpoints do not belong to StateGroupPool") # The free list, split by whether the group still carries content worth # something. Two containers rather than one queue because the two halves # want opposite orders and mixing them serves neither: @@ -207,16 +87,7 @@ def __init__( # rather than a second count because the depth is only ever one pass # more — see `release_pins`. self._deferred: set[int] = set() - # `copy` only. Seqs whose last forward left their state on a boundary - # worth keeping. `take_copies` turns each into a copy pair when the next - # batch is built, which is the latest moment the owner is still known to - # hold the group being duplicated. - self._checkpoint_pending: list = [] - # (src, dst) group pairs the next batch must copy before its forward. - # Both halves of the protocol feed this under `copy`: keeping a - # checkpoint copies the owner's state out, resuming from one copies it - # back in. - self._copies: list[tuple[int, int]] = [] + self._relocations: list[tuple[int, int]] = [] # `dropped` had no group to go to; `evicted` landed and was later # spent on an allocation. Counted apart because they read the same in a # hit rate and want opposite fixes — the first says the pool is too @@ -442,9 +313,7 @@ def resumable_hit( The fork test is what `min_fork_tokens` buys: resuming reads the checkpoint and writes a fresh group, and that forward has to leave the fresh group whole. A boundary too close to the end of the prompt fails - it and the scan keeps walking back. Under `copy` the resumer is handed - the bytes instead of reading across two groups, so `min_fork_tokens` is - 0 and the test is vacuous — one expression covers both. + it and the scan keeps walking back. Without this a hit hands the resumed forward a group freshly popped off the free list and it reads the previous occupant's state. @@ -457,13 +326,14 @@ def resumable_hit( return hit hbs = self.hash_block_size for i in range(hit - 1, -1, -1): - if not assume_checkpointed and block_hashes[i] not in self.hash_to_group: + checkpointed = block_hashes[i] in self.hash_to_group + if not assume_checkpointed and not checkpointed: continue if seq.num_tokens - (i + 1) * hbs >= self.min_fork_tokens: return i + 1 return 0 - def lookup(self, h: int) -> int: + def lookup_group(self, h: int) -> int: """Group holding the checkpoint for hash `h`, or -1.""" if not self.enabled: return -1 @@ -473,40 +343,13 @@ def lookup(self, h: int) -> int: def checkpoint(self, seq, boundary_blocks: int, h: int) -> None: """Keep `seq`'s state as of this boundary, filed under hash `h`. - Two mechanisms, chosen by the backend's `StateTransfer`: - - `fork` — the group cannot be shared while its owner still writes it, so - the owner moves to a fresh group and the old one, never written again, - becomes the checkpoint. The next forward reads it and fills the - replacement, which is the whole reason `min_fork_tokens` gates the - position. - - `copy` — the owner keeps writing where it is and a duplicate of its - state goes to the index instead. Only the intent is recorded here; the - destination group and the copy pair come from `take_copies` when the next - batch is built. Deferred because the bytes have to be moved by a forward, and - a checkpoint indexed before its bytes exist would hand a resuming - request whatever the destination happened to hold. - - `boundary_blocks` is unused here: a group is a single entry, not a span - of them. It is in the protocol for classes whose checkpoint is a run of - entries ending at the boundary. - - Best-effort under both: with no free group the seq simply keeps writing - its own and no checkpoint is taken. + The owner moves to a fresh group and the old group becomes read-only. """ - if not self.applies(seq): + if not self.applies(seq) or not self.transfer.forks: return old = seq.per_req_cache_group if old < 0: return - if self.transfer.copies: - if seq.pending_checkpoint == -1: - self._checkpoint_pending.append(seq) - # A later boundary supersedes an earlier one: the group holds the - # state as of the last forward, so only the last position is true. - seq.pending_checkpoint = h - return if not self.has_free(): self.checkpoints_dropped += 1 return @@ -522,40 +365,6 @@ def checkpoint(self, seq, boundary_blocks: int, h: int) -> None: self.pin(old, reader_is_next_batch=True) self.checkpoints_kept += 1 - def _commit_pending(self) -> None: - """Turn the last step's checkpoint intents into copy pairs. - - Each pending seq gets a destination group, which goes straight back on - the free list and into the index — the same capacity-neutral move - `checkpoint` makes under `fork`, covered by the same lazy eviction: - whoever pops the group next invalidates the hash on the way out. - - A seq preempted or finished in between carries no group any more and is - skipped, so nothing is ever indexed over state that is gone. That check - is only sound because this runs with the batch already decided — see - `take_copies`. - """ - if not self._checkpoint_pending: - return - copy_start = len(self._copies) - for seq in self._checkpoint_pending: - h, seq.pending_checkpoint = seq.pending_checkpoint, -1 - src = seq.per_req_cache_group - if h == -1 or src < 0: - continue - if self.lookup(h) >= 0: - continue - if not self.has_free(): - self.checkpoints_dropped += 1 - continue - dst = self.pop() - self._index(h, dst) - self._copies.append((src, dst)) - self.checkpoints_kept += 1 - self._checkpoint_pending.clear() - for _, dst in self._copies[copy_start:]: - self.release(dst) - def checkpoint_fates(self) -> dict[str, int]: """What became of the checkpoints the ladder asked this pool to keep.""" return { @@ -565,43 +374,13 @@ def checkpoint_fates(self) -> dict[str, int]: "checkpoints_orphaned": self.checkpoints_orphaned, } - def record_copy(self, src: int, dst: int) -> None: - """Schedule a state copy for the next batch's forward to issue.""" - self._copies.append((src, dst)) - - def take_copies(self) -> list[tuple[int, int]]: - """Every copy the batch now being built must issue before its forward. - - Called at the moment the batch is constructed, which is the whole point: - a checkpoint's source is the owner's live group, and it has to still be - that owner's when the copy runs. Committing earlier in the pass would - leave a window — an admission preempting that owner would return the - group to the free list, and the copy would then duplicate whatever the - next request wrote there into a group already indexed as a checkpoint. - Nothing runs between here and the batch, so the window is empty. - - A checkpoint therefore becomes visible one pass later than the step that - formed it, and this pass's own admissions get first claim on the free - list. Both are the right way round: admission is throughput, a - checkpoint is speculative reuse. - - The two kinds of pair cannot collide, so their order does not matter. A - resume source is a claimed or pinned checkpoint and a keeper source is a - live group; neither is on the free list, so `_commit_pending`'s `pop` - can return neither. - """ - self._commit_pending() - copies, self._copies = self._copies, [] - return copies + def record_relocation(self, src: int, dst: int) -> None: + """Schedule an Active Slot relocation for the next batch.""" + self._relocations.append((src, dst)) - def forget_pending(self, seq) -> None: - """Drop `seq`'s uncommitted checkpoint — its group is being released. - - The seq stays in `_checkpoint_pending` until the next commit, which - skips it on the cleared hash. Cheaper than removing it, and the list is - emptied every pass either way. - """ - seq.pending_checkpoint = -1 + def take_relocations(self) -> tuple[tuple[int, int], ...]: + relocations, self._relocations = self._relocations, [] + return tuple(relocations) def _index(self, h: int, group: int) -> None: """File `group` as the checkpoint for hash `h`. @@ -622,7 +401,7 @@ def _index(self, h: int, group: int) -> None: if prev != -1 and prev != group: self._set_hash(prev, -1) - def unindex(self, h: int) -> int: + def unindex(self, h: int) -> None: """Drop the checkpoint filed under `h`. The dual of `_index`. Called when the KV block of that hash leaves the block index. The two @@ -638,14 +417,17 @@ def unindex(self, h: int) -> int: it exact would have the state pool watch every block of every prefix, which costs more than the tail it would catch. - Returns the group freed, or -1. """ group = self.hash_to_group.get(h, -1) if group < 0: - return -1 + return self.invalidate(group) self.checkpoints_orphaned += 1 - return group + + def clear_index(self) -> None: + """Drop all checkpoint hashes, preserving only in-flight readers.""" + for group in list(self.hash_to_group.values()): + self.invalidate(group) def invalidate(self, group: int) -> None: """Drop `group`'s checkpoint. Called when the group is handed out.""" diff --git a/atom/model_engine/state_runtime.py b/atom/model_engine/state_runtime.py new file mode 100644 index 0000000000..f3684c2320 --- /dev/null +++ b/atom/model_engine/state_runtime.py @@ -0,0 +1,173 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +from collections.abc import Mapping +from dataclasses import dataclass, field +from math import inf + +from atom.model_engine.page_unit_checkpoint import ( + CheckpointRestoreOp, + CheckpointStoreOp, + PagedStateCheckpointSpec, +) + +FORK = "fork" +COPY = "copy" +NONE = "none" + + +@dataclass(frozen=True) +class StateTransfer: + """How a backend transfers one request's state to another slot.""" + + kind: str + fork_tokens: int = 0 + paged_layout_id: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.fork_tokens, int) or isinstance(self.fork_tokens, bool): + raise TypeError("fork_tokens must be an integer") + if self.kind == COPY: + if self.fork_tokens != 0: + raise ValueError("copy transfer cannot bind successor tokens") + if not isinstance(self.paged_layout_id, str) or not self.paged_layout_id: + raise ValueError("copy transfer requires a non-empty PAGE layout id") + return + if self.kind == FORK: + if self.fork_tokens <= 0: + raise ValueError("fork transfer requires positive fork_tokens") + if self.paged_layout_id is not None: + raise ValueError("fork transfer cannot declare a PAGE layout") + return + if self.kind == NONE: + if self.fork_tokens != 0 or self.paged_layout_id is not None: + raise ValueError("none transfer cannot carry tokens or a PAGE layout") + return + raise ValueError(f"unknown state transfer kind {self.kind!r}") + + @classmethod + def none(cls) -> "StateTransfer": + return cls(NONE) + + @classmethod + def fork(cls, tokens: int) -> "StateTransfer": + return cls(FORK, tokens) + + @classmethod + def copy(cls, layout_id: str) -> "StateTransfer": + return cls(COPY, paged_layout_id=layout_id) + + def to_wire(self) -> dict[str, str | int | None]: + return { + "kind": self.kind, + "fork_tokens": self.fork_tokens, + "paged_layout_id": self.paged_layout_id, + } + + @classmethod + def from_wire(cls, wire: object) -> "StateTransfer": + if not isinstance(wire, Mapping): + raise TypeError("state transfer capability must be a mapping") + expected = {"kind", "fork_tokens", "paged_layout_id"} + if set(wire) != expected: + raise ValueError( + "invalid state transfer capability fields: " + f"expected={sorted(expected)}, got={sorted(wire)}" + ) + return cls( + kind=wire["kind"], # type: ignore[arg-type] + fork_tokens=wire["fork_tokens"], # type: ignore[arg-type] + paged_layout_id=wire["paged_layout_id"], # type: ignore[arg-type] + ) + + @property + def copies(self) -> bool: + return self.kind == COPY + + @property + def forks(self) -> bool: + return self.kind == FORK + + @property + def successor_room(self) -> float: + return inf if self.kind == NONE else float(self.fork_tokens) + + +@dataclass(frozen=True) +class StateRuntime: + """Validated state transfer and optional checkpoint geometry.""" + + transfer: StateTransfer = field(default_factory=StateTransfer.none) + checkpoint_spec: PagedStateCheckpointSpec | None = None + + def __post_init__(self) -> None: + if not isinstance(self.transfer, StateTransfer): + raise TypeError("state runtime transfer must be a StateTransfer") + if self.checkpoint_spec is not None and not isinstance( + self.checkpoint_spec, PagedStateCheckpointSpec + ): + raise TypeError( + "state runtime checkpoint_spec must be a PagedStateCheckpointSpec" + ) + if self.transfer.copies: + if self.checkpoint_spec is None: + raise ValueError( + "StateTransfer.copy(layout_id) requires a PAGE checkpoint spec" + ) + if self.transfer.paged_layout_id != self.checkpoint_spec.layout_id: + raise ValueError( + "state runtime PAGE layout mismatch: " + f"transfer={self.transfer.paged_layout_id!r}, " + f"spec={self.checkpoint_spec.layout_id!r}" + ) + elif self.checkpoint_spec is not None: + raise ValueError( + f"StateTransfer.{self.transfer.kind} cannot carry a PAGE checkpoint spec" + ) + + def to_wire(self) -> dict[str, object]: + return { + "transfer": self.transfer.to_wire(), + "checkpoint_spec": ( + None if self.checkpoint_spec is None else self.checkpoint_spec.to_wire() + ), + } + + @classmethod + def from_wire(cls, wire: object) -> "StateRuntime": + if not isinstance(wire, Mapping): + raise TypeError("state runtime must be a mapping") + expected = {"transfer", "checkpoint_spec"} + if set(wire) != expected: + raise ValueError( + "invalid state runtime fields: " + f"expected={sorted(expected)}, got={sorted(wire)}" + ) + checkpoint_wire = wire["checkpoint_spec"] + checkpoint_spec = ( + None + if checkpoint_wire is None + else PagedStateCheckpointSpec.from_wire(checkpoint_wire) + ) + return cls( + transfer=StateTransfer.from_wire(wire["transfer"]), + checkpoint_spec=checkpoint_spec, + ) + + +DEFAULT_STATE_RUNTIME = StateRuntime() + + +@dataclass(frozen=True) +class StateMaintenanceOps: + """State movement drained once before a model batch.""" + + relocations: tuple[tuple[int, int], ...] = () + checkpoint_stores: tuple[CheckpointStoreOp, ...] = () + checkpoint_restores: tuple[CheckpointRestoreOp, ...] = () + + @property + def empty(self) -> bool: + return not ( + self.relocations or self.checkpoint_stores or self.checkpoint_restores + ) diff --git a/atom/model_ops/attentions/backends.py b/atom/model_ops/attentions/backends.py index 0109ae68a1..2ab2ddfa68 100644 --- a/atom/model_ops/attentions/backends.py +++ b/atom/model_ops/attentions/backends.py @@ -3,6 +3,7 @@ import logging from abc import ABC, abstractmethod +from collections.abc import Sequence from typing import TYPE_CHECKING, Any, Generic, Optional, TypeVar if TYPE_CHECKING: @@ -14,8 +15,12 @@ from torch import nn from atom.distributed.dcp_utils import get_dcp_rank, get_dcp_world_size +from atom.model_engine.page_unit_checkpoint import ( + CheckpointRestoreOp, + CheckpointStoreOp, +) from atom.model_engine.scheduler import ScheduledBatch -from atom.model_engine.state_pool import StateTransfer +from atom.model_engine.state_runtime import StateTransfer from atom.model_ops.attention_mla import MLAModules from atom.model_ops.attentions.sub_pool_spec import SubPoolSpec from atom.utils import CpuGpuBuffer @@ -139,54 +144,27 @@ def allocate_per_req_cache(self, entries: dict[str, int]) -> dict[str, object]: return {} def state_transfer(self) -> StateTransfer: - """How this backend hands one request's state to another group. - - A checkpoint is a second group holding the state as of some boundary, so - every backend with per-request state has to say how one gets there. - There are three answers and `StateGroupPool` runs whichever it is told: - - `StateTransfer.fork(n)` — the state rolls and is not one range to - duplicate, so the old group goes to the index and the request takes a - fresh one, reading the old and writing the new for exactly one forward. - That forward has to leave the new group self-contained (a single read - index cannot span both), which takes `n` *committed* tokens. - `BlockManager` walks a checkpoint/hit point back to the previous block - boundary until it fits. - - `StateTransfer.copy()` — one request's state is a contiguous byte range, - so the index gets a duplicate and the owner is left alone. No forward is - bound and no boundary is disqualified for lack of room, which is what - makes a decode boundary checkpointable at all: a decode step commits - `1 + accepted_drafts` tokens and acceptance is not knowable when the - checkpoint has to be decided. The backend must implement - `copy_state_entries`. - - `StateTransfer.none()` (default) — no per-request state, or none that can - be handed over; the checkpoint index stays empty and prefix hits shrink - to 0 for its models. - """ + """Declare this backend's per-request state checkpoint capability.""" return StateTransfer.none() - def copy_state_entries(self, pairs: list[tuple[int, int]]) -> None: - """Copy each `(src, dst)` group's whole per-request state, src → dst. - - Issued by `build` before the forward, on the compute stream, so a copy - lands after the forward that produced its source and before the one that - consumes its destination. - - Owed by every backend that declares a state pool, not just the ones - declaring `StateTransfer.copy()`. Two callers want it and only the first - is about checkpointing: a copy-transfer class duplicates a group to keep - a checkpoint, and *any* class has to be able to hand a group's bytes to - a different group index when the pool's boundary moves past the one it - is sitting on. The second is a byte move regardless of how the class - checkpoints, so a fork-transfer backend owes this too. - """ + def relocate_state_slots(self, pairs: Sequence[tuple[int, int]]) -> None: + """Move live state between contiguous Active Slots.""" raise NotImplementedError( f"{type(self).__name__} owns per-request state but does not " - "implement copy_state_entries" + "implement relocate_state_slots" ) + def execute_paged_state_copies( + self, + store_ops: Sequence[CheckpointStoreOp], + restore_ops: Sequence[CheckpointRestoreOp], + ) -> None: + """Copy checkpoints between Active Slots and arbitrary PAGEs.""" + if store_ops or restore_ops: + raise NotImplementedError( + f"{type(self).__name__} does not implement PAGE-backed state copy" + ) + def get_kv_transfer_tensors(self) -> "KVTransferTensors | None": """Return RDMA transfer regions for PD disaggregation. @@ -517,14 +495,14 @@ def _attach_tbo_prefill_cpu_lens( ) def build(self, batch: ScheduledBatch, bs: int): - # State checkpoints the scheduler decided on ride the batch as group - # pairs and are copied here, on the compute stream, before the forward. - # This is the one place every path — prefill, decode, dummy, DP-sync, PP - # microbatch, TBO — passes through exactly once per batch, which is what - # makes "each copy is issued once per rank" true by construction rather - # than by inspection of every prepare_* variant. - if batch.state_copy_pairs: - self.copy_state_entries(batch.state_copy_pairs) + # Run state maintenance on the compute stream before the forward. + state_ops = batch.state_maintenance_ops + if state_ops.relocations: + self.relocate_state_slots(state_ops.relocations) + if state_ops.checkpoint_stores or state_ops.checkpoint_restores: + self.execute_paged_state_copies( + state_ops.checkpoint_stores, state_ops.checkpoint_restores + ) is_prefill = batch.total_tokens_num_prefill > 0 if is_prefill: return self.prepare_prefill(batch) diff --git a/atom/model_ops/attentions/deepseek_v4_attn.py b/atom/model_ops/attentions/deepseek_v4_attn.py index 6190aa865a..3da65c0fdd 100644 --- a/atom/model_ops/attentions/deepseek_v4_attn.py +++ b/atom/model_ops/attentions/deepseek_v4_attn.py @@ -40,6 +40,7 @@ import logging import math import os +from collections.abc import Sequence from dataclasses import dataclass from typing import Any, cast @@ -58,8 +59,12 @@ pcp_round_robin_query_indices, ) from atom.model_engine.kv_block import STATE_SLOT_CLASS +from atom.model_engine.page_unit_checkpoint import ( + CheckpointRestoreOp, + CheckpointStoreOp, +) from atom.model_engine.scheduler import ScheduledBatch -from atom.model_engine.state_pool import StateTransfer +from atom.model_engine.state_runtime import StateTransfer from atom.model_ops.attentions.backends import ( AttentionBackend, AttentionMetadataBuilder, @@ -730,12 +735,7 @@ def _state_fields(self) -> list[StateField]: because a plane is one width and a field is the one thing in this layout that is priced in bytes. - Putting it here rather than in a plane of its own is what makes it - travel: `copy_state_entries` copies a whole slot, so a checkpoint - carries the window with the state. A private plane would not be - copied, and a request resuming a cached prefix would draft against - whatever the slot's previous occupant left — the same shape of bug - #1417 was, minus the correctness half, since drafts are verified. + Keeping it in the slot makes checkpoint scatter carry the window too. Field order is the wire order of a whole entry, so it is also the order a PD transfer or a checkpoint sees the bytes in. @@ -767,51 +767,22 @@ def _state_fields(self) -> list[StateField]: return fields def state_transfer(self) -> StateTransfer: - """A copy: one request's compressor state is one contiguous entry. - - `StateArena` already lays a request's whole compressor state out as one - byte range (`entry(i)`), which is what makes the duplicate a single - `copy_` — see `copy_state_entries`. - - A fork would also work on a prompt and would move no bytes, but it binds - the forward after the checkpoint: that forward has to leave the fresh - group self-contained, which takes `K - ratio` = 4 *committed* tokens for - the overlapping CSA ring (0 for HCA). A decode step commits - `1 + accepted_drafts`, and acceptance is not knowable when the - checkpoint has to be decided — nor recoverable afterwards, since by then - the state is split across two groups and a single read index spans - neither. So a fork can never checkpoint a decode boundary, and this - model's reuse is multi-turn, where the boundary worth keeping is exactly - the one generation ends on. - - Arithmetic for both numbers, replayed from `compress_plan.py`: - `/app/logs_claude/verify_v4_min_fork.py`. - """ - return StateTransfer.copy() - - def copy_state_entries(self, pairs: list[tuple[int, int]]) -> None: - """Duplicate a request's whole per-request state: compressor + windows. - - One range per plane, and a plane is all there is: a slot holds the - compressor state and then every layer's windows, contiguously, so the - two halves of a request's state copy as one slice. Copying the state - whole rather than just the rows a resumer reads (the CSA ring's - trailing `K - ratio` = 4, HCA's none) is what makes that true — those - rows are scattered, and picking them out would cost 84 strided copies - against this one. - - **The window half is what makes the private ring safe.** A per-request - ring is exactly what #1417 removed, because a request resuming a cached - prefix had never written that prefix into its own. Reinstating it is only - correct because the checkpoint carries the window across — drop this and - the bug returns, silently, as garbage attention over the reused prefix. - - That is also why a window whose dtype differs from the pool's is a state - field rather than a plane of its own: a plane of its own would sit - outside every slot and so outside this copy, and a resumed request would - draft against the slot's previous occupant. Verification would keep the - output right and the acceptance rate would just quietly collapse. - """ + """Declare PAGE-copy checkpoints with the versioned DSV4 layout.""" + ratios = ",".join(str(r) for r in self._geometry_ratios()) + layout_id = ( + "dsv4-paged-state-v1" + f":block={self.block_size}:ring={self.win_with_spec}" + f":dims={self.head_dim},{self.rope_head_dim},{self.index_head_dim}" + f":state={self.csa_main_state_shape},{self.csa_idx_state_shape}," + f"{self.hca_main_state_shape}" + f":main={'fp8-2buff' if self._kv_fp8 else 'bf16'}" + f":index={'fp4' if self._indexer_fp4 else 'fp8'}" + f":ratios={ratios}" + ) + return StateTransfer.copy(layout_id) + + def relocate_state_slots(self, pairs: Sequence[tuple[int, int]]) -> None: + """Relocate a request's whole Active Slot.""" views = self._slot_views() dsts, srcs = [], [] for src, dst in pairs: @@ -820,6 +791,88 @@ def copy_state_entries(self, pairs: list[tuple[int, int]]) -> None: if dsts: torch._foreach_copy_(dsts, srcs) + def execute_paged_state_copies( + self, + store_ops: Sequence[CheckpointStoreOp], + restore_ops: Sequence[CheckpointRestoreOp], + ) -> None: + """Copy raw checkpoint bytes between slots and non-contiguous PAGEs.""" + from atom.model_ops.attentions.paged_state_copy import ( + launch_copy_spans, + plan_segmented_copy, + ) + + spans = [] + device = self._kv_planes()[0].device + for op in store_ops: + self._validate_paged_state_op(op) + src = self._active_slot_segments(op.src_slot) + dst = self._page_unit_segments(op.unit_ids) + spans.extend(plan_segmented_copy(src, dst, op.total_bytes)) + for op in restore_ops: + self._validate_paged_state_op(op) + src = self._page_unit_segments(op.unit_ids) + dst = self._active_slot_segments(op.dst_slot) + spans.extend(plan_segmented_copy(src, dst, op.total_bytes)) + launch_copy_spans(spans, device) + + def _validate_paged_state_op( + self, op: CheckpointStoreOp | CheckpointRestoreOp + ) -> None: + spec = self.model_runner.state_runtime.checkpoint_spec + if spec is None: + raise RuntimeError("DSV4 PAGE/state checkpoint spec is missing") + if op.layout_id != spec.layout_id: + raise RuntimeError( + f"state checkpoint layout mismatch: {op.layout_id!r} != " + f"{spec.layout_id!r}" + ) + if op.total_bytes != spec.slot_bytes: + raise RuntimeError( + f"state checkpoint size mismatch: op={op.total_bytes}, " + f"active_slot={spec.slot_bytes}" + ) + if len(op.unit_ids) != spec.units_per_checkpoint: + raise RuntimeError( + "state checkpoint PAGE-unit geometry does not match this worker" + ) + num_blocks = self.model_runner.num_physical_kvcache_blocks + if any(unit_id < 0 or unit_id >= num_blocks for unit_id in op.unit_ids): + raise RuntimeError("state checkpoint PAGE unit is out of range") + + def _active_slot_segments(self, group: int): + from atom.model_ops.attentions.paged_state_copy import tensor_segment + + return [tensor_segment(view) for view in self._slot_views()[group]] + + def _one_page_unit_segments(self, block_id: int): + """All physical regions owned by one logical DSV4 PAGE block id.""" + from atom.model_ops.attentions.paged_state_copy import tensor_segment + + runner = self.model_runner + geo = self.pool_geometry + start = block_id * geo.envelope_rows + stop = start + geo.envelope_rows + segments = [tensor_segment(plane[start:stop]) for plane in self._kv_planes()] + segments.extend( + tensor_segment(runner.v4_csa_idx_kv[layer, block_id]) + for layer in range(len(self.csa_layers)) + ) + if self._indexer_fp4: + segments.extend( + tensor_segment(runner.v4_csa_idx_kv_scale[layer, block_id]) + for layer in range(len(self.csa_layers)) + ) + return segments + + def _page_unit_segments(self, unit_ids): + segments = [] + for block_id in unit_ids: + if not 0 <= block_id < self.model_runner.num_physical_kvcache_blocks: + raise RuntimeError(f"PAGE unit id {block_id} is out of range") + segments.extend(self._one_page_unit_segments(block_id)) + return segments + def _slot_views(self) -> list[list[torch.Tensor]]: """Per-group views of that request's whole slot in each plane. @@ -1029,6 +1082,30 @@ def allocate_per_req_cache(self, entries: dict[str, int]) -> dict[str, object]: geo = self.pool_geometry.with_capacity(num_blocks, num_slots) self.pool_geometry = geo + actual_page_bytes = ( + sum(geo.block_bytes(width) for width in self._plane_row_widths()) + + self._indexer_block_bytes() + ) + actual_slot_bytes = sum( + geo.slot_bytes(width) for width in self._plane_row_widths() + ) + state_runtime = self.model_runner.state_runtime + checkpoint_spec = state_runtime.checkpoint_spec + if checkpoint_spec is None: + raise RuntimeError("DSV4 PAGE/state checkpoint sizing spec is missing") + layout_id = state_runtime.transfer.paged_layout_id + if ( + actual_page_bytes != checkpoint_spec.page_unit_bytes + or actual_slot_bytes != checkpoint_spec.slot_bytes + or layout_id != checkpoint_spec.layout_id + ): + raise RuntimeError( + "DSV4 PAGE/state checkpoint geometry differs from sizing: " + f"page={actual_page_bytes}/{checkpoint_spec.page_unit_bytes}, " + f"slot={actual_slot_bytes}/{checkpoint_spec.slot_bytes}, " + f"layout={layout_id!r}/{checkpoint_spec.layout_id!r}" + ) + row_widths = self._plane_row_widths() offsets, total_bytes = plan_regions([geo.plane_bytes(w) for w in row_widths]) diff --git a/atom/model_ops/attentions/gdn_attn.py b/atom/model_ops/attentions/gdn_attn.py index 6dda2d90c1..aa06a3cc2a 100644 --- a/atom/model_ops/attentions/gdn_attn.py +++ b/atom/model_ops/attentions/gdn_attn.py @@ -2,6 +2,7 @@ # Copyright (C) 2024-2025, Advanced Micro Devices, Inc. All rights reserved. import math +from collections.abc import Sequence from dataclasses import dataclass import numpy as np @@ -10,7 +11,7 @@ from atom.model_engine.kv_block import STATE_SLOT_CLASS from atom.model_engine.scheduler import ScheduledBatch -from atom.model_engine.state_pool import StateTransfer +from atom.model_engine.state_runtime import StateTransfer from atom.model_ops.attention_gdn import GatedDeltaNet from atom.utils import CpuGpuBuffer from atom.utils.forward_context import AttentionMetaData, Context @@ -257,23 +258,7 @@ def _state_shape_for_runner(self) -> tuple[tuple[int, ...], tuple[int, ...]]: ) def state_transfer(self) -> StateTransfer: - """A fork whose successor forward need only carry one token. - - Both halves of the GDN state come out of a forward self-contained at any - length. The recurrent state is rewritten whole, and every write path in - `causal_conv1d` stores the full `state_len` window to the output slot — - the short-chunk paths get there by loading the previous window from the - *input* slot, shifting left and appending x — so the new group stops - depending on the old one the moment the forward returns. - - Reading the state layout alone suggests `conv_kernel_dim - 1` instead, - on the theory that a shorter forward leaves the new group holding a - window the old group still owns part of. The kernel closes that gap. - - A fork rather than a copy because the state is two per-family tensors - rather than one contiguous entry, so there is no single range to - duplicate — and at one token the fork binds almost nothing anyway. - """ + """Declare one-token fork checkpoint support for recurrent state.""" return StateTransfer.fork(1) def state_spec(self) -> SubPoolSpec: @@ -315,8 +300,8 @@ def allocate_per_req_cache( ), } - def copy_state_entries(self, pairs: list[tuple[int, int]]) -> None: - """Duplicate a group's whole GDN state, both families, all layers. + def relocate_state_slots(self, pairs: Sequence[tuple[int, int]]) -> None: + """Relocate a live GDN group between logical Active Slot spans. A group is `1 + num_spec` consecutive slots — the extra ones hold the per-draft states a rejected speculation rolls back to — so a group moves @@ -324,9 +309,7 @@ def copy_state_entries(self, pairs: list[tuple[int, int]]) -> None: GDN checkpoints by forking, not by copying, so this is not on the checkpoint path: it exists because moving the pool's boundary has to be - able to relocate a group that is in the way, and relocation is a byte - move whatever mechanism the class uses to checkpoint. A backend - declaring `StateTransfer.fork` therefore still owes this method. + able to relocate a group that is in the way. Both caches are layer-major with the slot as the second axis, so a group's rows are strided rather than contiguous and there is no single diff --git a/atom/model_ops/attentions/paged_state_copy.py b/atom/model_ops/attentions/paged_state_copy.py new file mode 100644 index 0000000000..3ccab5497e --- /dev/null +++ b/atom/model_ops/attentions/paged_state_copy.py @@ -0,0 +1,138 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""Descriptor-driven bitwise copy between segmented GPU byte streams.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +try: + import triton + import triton.language as tl +except ModuleNotFoundError: + triton = None + tl = None + +_TILE_BYTES = 4096 + + +@dataclass(frozen=True) +class ByteSegment: + ptr: int + num_bytes: int + + +@dataclass(frozen=True) +class CopySpan: + src_ptr: int + dst_ptr: int + num_bytes: int + + +def tensor_segment(tensor: torch.Tensor) -> ByteSegment: + """Describe a contiguous tensor view as raw bytes without converting it.""" + if not tensor.is_contiguous(): + raise ValueError("paged state copy segments must be contiguous") + return ByteSegment(int(tensor.data_ptr()), tensor.numel() * tensor.element_size()) + + +def plan_segmented_copy( + src: list[ByteSegment], + dst: list[ByteSegment], + total_bytes: int, +) -> list[CopySpan]: + """Intersect two ordered byte streams into physical copy spans.""" + total_bytes = int(total_bytes) + if total_bytes < 0: + raise ValueError("copy length must be non-negative") + if sum(s.num_bytes for s in src) < total_bytes: + raise ValueError("source segmented stream is shorter than the copy") + if sum(s.num_bytes for s in dst) < total_bytes: + raise ValueError("destination segmented stream is shorter than the copy") + if any(s.num_bytes <= 0 for s in src + dst): + raise ValueError("segmented streams cannot contain empty segments") + if total_bytes == 0: + return [] + + spans: list[CopySpan] = [] + src_i = dst_i = 0 + src_off = dst_off = 0 + remaining = total_bytes + while remaining: + src_left = src[src_i].num_bytes - src_off + dst_left = dst[dst_i].num_bytes - dst_off + nbytes = min(src_left, dst_left, remaining) + spans.append( + CopySpan( + src[src_i].ptr + src_off, + dst[dst_i].ptr + dst_off, + nbytes, + ) + ) + remaining -= nbytes + src_off += nbytes + dst_off += nbytes + if src_off == src[src_i].num_bytes: + src_i += 1 + src_off = 0 + if dst_off == dst[dst_i].num_bytes: + dst_i += 1 + dst_off = 0 + return spans + + +if triton is not None: + + @triton.jit + def _copy_tiles_kernel( + src_ptrs, + dst_ptrs, + valid_bytes, + TILE_BYTES: tl.constexpr, + ): + tile = tl.program_id(0) + offsets = tl.arange(0, TILE_BYTES) + valid = tl.load(valid_bytes + tile) + mask = offsets < valid + src_addr = tl.load(src_ptrs + tile).to(tl.int64) + dst_addr = tl.load(dst_ptrs + tile).to(tl.int64) + src = (src_addr + offsets).to(tl.pointer_type(tl.uint8)) + dst = (dst_addr + offsets).to(tl.pointer_type(tl.uint8)) + value = tl.load(src, mask=mask) + tl.store(dst, value, mask=mask) + +else: + _copy_tiles_kernel = None + + +def launch_copy_spans(spans: list[CopySpan], device: torch.device) -> None: + """Copy all spans with one descriptor-driven Triton launch.""" + if not spans: + return + if _copy_tiles_kernel is None: + raise RuntimeError("paged state copy requires Triton") + src_ptrs: list[int] = [] + dst_ptrs: list[int] = [] + valid_bytes: list[int] = [] + for span in spans: + offset = 0 + while offset < span.num_bytes: + nbytes = min(_TILE_BYTES, span.num_bytes - offset) + src_ptrs.append(span.src_ptr + offset) + dst_ptrs.append(span.dst_ptr + offset) + valid_bytes.append(nbytes) + offset += nbytes + + src_t = torch.tensor(src_ptrs, dtype=torch.int64, device=device) + dst_t = torch.tensor(dst_ptrs, dtype=torch.int64, device=device) + valid_t = torch.tensor(valid_bytes, dtype=torch.int32, device=device) + _copy_tiles_kernel[(len(src_ptrs),)]( + src_t, + dst_t, + valid_t, + TILE_BYTES=_TILE_BYTES, + num_warps=8, + ) diff --git a/atom/model_ops/attentions/sub_pool_spec.py b/atom/model_ops/attentions/sub_pool_spec.py index d63179ead6..2ec950bdae 100644 --- a/atom/model_ops/attentions/sub_pool_spec.py +++ b/atom/model_ops/attentions/sub_pool_spec.py @@ -45,8 +45,6 @@ class — and the backend that declares the spec imports it from there. The one from dataclasses import dataclass, replace from enum import Enum -from atom.utils import envs - class Pool(Enum): """Which budget region an entry class draws from.""" @@ -96,8 +94,6 @@ def state_pool( extra_entries: int = 0, ) -> SubPoolSpec: """A per-request state entry class.""" - if envs.is_set("STATE_CKPT_EXTRA_ENTRIES"): - extra_entries = int(envs.STATE_CKPT_EXTRA_ENTRIES) return SubPoolSpec(Pool.STATE, name, entry_bytes, entries_per_req, extra_entries) diff --git a/atom/utils/envs.py b/atom/utils/envs.py index 4d1750bb3e..45491362ee 100644 --- a/atom/utils/envs.py +++ b/atom/utils/envs.py @@ -45,9 +45,6 @@ # ATOM remaps the SGLang world into internal TP x PCP groups. # 0 means unset. "ATOM_SGLANG_PCP_SIZE": lambda: int(os.getenv("ATOM_SGLANG_PCP_SIZE", "0") or "0"), - "STATE_CKPT_EXTRA_ENTRIES": lambda: int( - os.getenv("STATE_CKPT_EXTRA_ENTRIES", "0") or "0" - ), # --- Compilation & Execution --- "ATOM_USE_TRITON_GEMM": lambda: os.getenv("ATOM_USE_TRITON_GEMM", "0") == "1", "ATOM_FP8_BLOCKSCALE_USE_E8M0_SCALE": lambda: ( diff --git a/docs/scheduling_kv_cache_guide.md b/docs/scheduling_kv_cache_guide.md index 6c60027d7e..99af8cf5a3 100644 --- a/docs/scheduling_kv_cache_guide.md +++ b/docs/scheduling_kv_cache_guide.md @@ -218,7 +218,7 @@ Methods: ```python class BlockManager: - def __init__(self, config: Config): + def __init__(self, config: Config, *, state_runtime: StateRuntime = DEFAULT_STATE_RUNTIME): block_size = config.kv_cache_block_size # Tokens per block (default 16) num_blocks = config.num_kvcache_blocks # Total blocks in pool self.block_size = block_size @@ -237,11 +237,23 @@ class BlockManager: state_entries = int(pool_entries.get(STATE_SLOT_CLASS, 0)) state_per_req = int(pool_per_req.get(STATE_SLOT_CLASS, 1)) or 1 self.num_per_req_cache_groups = state_entries // state_per_req + checkpoint_spec = state_runtime.checkpoint_spec + self.paged_state_checkpoints = ( + None + if checkpoint_spec is None + else PagedStateCheckpointCoordinator( + self.kv, + checkpoint_spec, + enabled=self.enable_prefix_caching + and self.num_per_req_cache_groups > 0, + ) + ) self.state = StateGroupPool( self.num_per_req_cache_groups, - transfer=StateTransfer.from_config( - getattr(config, "state_transfer_kind", "none") or "none", - int(getattr(config, "state_fork_tokens", 0) or 0), + transfer=( + StateTransfer.none() + if self.paged_state_checkpoints is not None + else state_runtime.transfer ), hash_block_size=self.hash_block_size, enabled=self.enable_prefix_caching, @@ -264,7 +276,7 @@ where `q = pos % win_with_spec`. The layer term is not in it: a layer's view is It was a content-addressed block pool until this change, which is worth writing down because the choice is not obvious and it is not permanent. **The question is where reuse comes from: a block pool reuses by sharing rows, a ring reuses by copying them.** Everything else follows. -Sharing rows was the only mechanism available before per-request state could be checkpointed by copying (`StateGroupPool` under `StateTransfer.copy()`). A private ring at that time meant a request resuming someone else's cached prefix had never written that prefix into its own ring and read stale rows — issue #1417, which is exactly what replaced the ring with a pool. The ring is back only because `copy_state_entries` now carries the window across with the compressor state. **Reverting the addressing without that copy reintroduces #1417 silently**, so the two are one change, not two. +Sharing rows was the only mechanism available before per-request state could be checkpointed by copying (`PagedStateCheckpointCoordinator` under `StateTransfer.copy(layout_id)`). A private ring at that time meant a request resuming someone else's cached prefix had never written that prefix into its own ring and read stale rows — issue #1417, which is exactly what replaced the ring with a pool. The ring is back only because `execute_paged_state_copies` now gathers the saved window and compressor state into the resumer's Active Slot. **Reverting the addressing without that copy reintroduces #1417 silently**, so the two are one change, not two. What the ring buys: @@ -283,7 +295,7 @@ What it costs: **When the trade reverses.** The memory win is entirely the `window / block_size` ratio: at V4's 128/256 a ring is 4× smaller, but at a 2048-token window a block pool needs `ceil(2048/256)+1 = 9` blocks = 2304 tokens for 2048, and the ring saves almost nothing while keeping all of its aliasing invariants. Note also that sharing rows was worth less than it looks: a resuming request shares only the trailing window and starts writing its own rows immediately, so a block pool never held one window for N requests either. **If V4's window ever grows past its block size, revisit this.** **Per-Request Cache Pools (Stateful-Attention Models):** For models whose attention type maintains per-request state outside the paged KV pool (GDN: Qwen3-Next, Qwen3.5, Kimi-Linear; DeepSeek-V4's compressor ring): -- `state` — a [`StateGroupPool`](../atom/model_engine/state_pool.py), owning both the free list of group indices (0 to `num_per_req_cache_groups - 1`) and the content index over them. Each group is one request's worth: `entries_per_req` contiguous tensor slot indices (1 for a single committed state, `1 + num_speculative_tokens` where a rollback slot per speculated token is kept). See **State checkpoints** below. +- `state` — a [`StateGroupPool`](../atom/model_engine/state_pool.py), owning the Active Slot free list and the fork-checkpoint index. PAGE-copy checkpoints are owned separately by [`PagedStateCheckpointCoordinator`](../atom/model_engine/page_unit_checkpoint.py). Each group is one request's worth: `entries_per_req` contiguous tensor slot indices (1 for a single committed state, `1 + num_speculative_tokens` where a rollback slot per speculated token is kept). See **State checkpoints** below. - `num_per_req_cache_groups` — total capacity, so callers can tell "all slots busy" (transient) from "no slots were ever created" (permanent). The state class costs no paged blocks at admission time: sizing reserves every STATE class's floor before the paged class is sized (see [`sub_pool_spec.py`](../atom/model_ops/attentions/sub_pool_spec.py)), so a sequence only needs a free slot index. Because that floor is exactly `max_num_seqs` requests' worth, the slot pool never binds before `max_num_seqs` does. @@ -303,7 +315,7 @@ Run to a fixpoint rather than `min()`-ed or chained: the answer has to satisfy e That number, `successor_room`, is mutability quantified. A rolling state (GDN recurrence) is still being written by its owner and is not one range to duplicate, so keeping it means handing the group over and taking a fresh one — and the next forward has to refill the replacement, which is `min_fork_tokens` of it. An immutable entry, or one that can simply be copied, needs no hand-over and no successor, i.e. `0`. `inf` means the class cannot be checkpointed at all — it would gate hits and never keep one. No class reports it today; `StateTransfer.none()` decodes to it, so a backend with no transferable state lands there rather than being special-cased. -**Checkpoints cost no capacity.** A checkpoint *is* a group sitting on the free list with its content intact, indexed by the content hash of the last block it covers — the same lazy-eviction model the block pool uses, where hand-out (`StateGroupPool.pop`), not free, is the eviction event. The pool therefore never holds a group back, and under full concurrency the checkpoint set drains on its own. +**Checkpoint capacity follows the transfer kind.** A fork checkpoint *is* a group sitting on the free list with its content intact, indexed by the content hash of the last block it covers — the same lazy-eviction model the block pool uses, where hand-out (`StateGroupPool.pop`), not free, is the eviction event. The pool therefore never holds a group back, and under full concurrency the fork checkpoint set drains on its own. A copy checkpoint does not occupy an Active Slot: it owns an ordered set of arbitrary PAGE units and is reclaimed only as a whole record. Active Slots remain reserved for resident requests. **The free list is two halves.** Groups carrying nothing sit in one container ordered by index; groups carrying a checkpoint sit in another ordered least-recently-used. `pop` always drains the first before touching the second, so a checkpoint can only be spent once there is nothing free left to take — a single release-ordered queue cannot express that, because a checkpoint handed back before a never-used group sits ahead of it and is spent first. Reuse counts as use: `claim` leaves the hash in place, so a resumed checkpoint returns through `release` to the LRU tail. @@ -311,20 +323,22 @@ Index order in the vacant half is not a fairness choice. Allocating lowest-first **Two ways to keep one.** How a group reaches the index is the backend's `StateTransfer`, declared by `AttentionMetadataBuilder.state_transfer()`, and it decides *where* that backend may checkpoint. +After pool sizing, the runner combines that capability with the optional `PagedStateCheckpointSpec` in one validated `StateRuntime`. COPY requires a spec with the same versioned layout id; FORK and NONE forbid one. The nested runtime wire payload is reconstructed once in `EngineCore` and passed explicitly through the scheduler to `BlockManager`, so neither `Config` nor downstream constructors can observe an invalid transfer/spec combination. + *Fork* (`StateTransfer.fork(n)`, GDN). At a rung the request hands its group to the index and takes a fresh one; for exactly one forward it then reads the handed-over group and writes the new one (`non_spec_state_indices_in_tensor` / `non_spec_state_indices_tensor`). A checkpointed group is never written again, which is what makes it safe to share. Resuming is the same move in reverse. The cost is that the *next* forward is bound: it has to leave the replacement self-contained, which takes `n` committed tokens. -*Copy* (`StateTransfer.copy()`, DeepSeek-V4). The state is one contiguous entry ([`StateArena`](../atom/model_ops/attentions/state_arena.py)), so a duplicate goes to the index and the owner is not disturbed. Nothing downstream has to cooperate, which is what makes a decode boundary checkpointable at all — see below. The bytes still need a forward to move them, so `checkpoint` only records the intent; `StateGroupPool.take_copies`, at the moment the next batch is built, takes a destination group and emits a `(src, dst)` pair — that late because the source is the owner's live group, and an earlier commit would leave a window in which an admission preempts that owner and the copy duplicates the next request's state instead. `ScheduledBatch.state_copy_pairs` carries the pairs and `AttentionMetadataBuilder.build` issues them on the compute stream before the forward — one place every path passes through exactly once per batch. Deferring the index entry until the copy is scheduled is what stops a resuming request claiming a checkpoint whose bytes do not exist yet. +*Copy* (`StateTransfer.copy(layout_id)`, DeepSeek-V4). The resident request keeps one contiguous Active Slot ([`StateArena`](../atom/model_ops/attentions/state_arena.py)); an immutable checkpoint scatters that slot's canonical byte stream across an ordered set of arbitrary PAGE units, so the owner is not disturbed and checkpoint allocation does not require a large contiguous extent. Nothing downstream has to cooperate, which is what makes a decode boundary checkpointable at all — see below. The bytes still need a forward to move them, so `checkpoint` records the intent and keeps the record invisible in `COPYING` state. When the next real batch is built, `BlockManager.take_state_maintenance_ops` is the only drain: it returns one typed `StateMaintenanceOps` containing Active Slot relocations, checkpoint stores and checkpoint restores. `ScheduledBatch.state_maintenance_ops` carries that bundle, and `AttentionMetadataBuilder.build` issues every operation on the compute stream before the forward — one place every execution path passes through exactly once per batch. An empty batch does not drain the bundle. Publishing the checkpoint only after its store has ridden a real batch stops a resumer claiming bytes that do not exist yet. -Under either mechanism, when no second group is free the request adopts the checkpoint instead, spending it rather than sharing it. +Under fork, when no second group is free the request can adopt the checkpoint group, spending it rather than sharing it. Copy never adopts PAGE fragments as an Active Slot: a hit first obtains a complete contiguous Active Slot and then gathers the checkpoint into it; without a free Active Slot, admission waits. -**Where checkpoints land.** One ladder for every state class: a rung every `--state-checkpoint-interval-tokens` (default 8192) of context. Whether a class takes a given rung comes down to one comparison — how many tokens the *next* forward carries, against that class's `successor_room`. `BlockManager.checkpointers_at` takes the first as an argument (prefill passes what is left of the prompt, decode passes one token) and returns the classes that qualify; `checkpoint_limit` is the same rule solved for prefill's last qualifying rung and `checkpoint_cut` turns it into a chunk boundary, which the scheduler needs up front (`_finalize_prefill_chunk`). Everything else follows from the one comparison: GDN's `fork(1)` always qualifies — its `causal_conv1d` write paths all store the full window to the output slot — V4's `copy()` reports 0 and so does too, and a rolling class needing a long hand-over simply never qualifies mid-generation. A backend with no transferable state at all declares `StateTransfer.none()`, which is `inf` on this scale; it is a separate kind rather than a token count precisely because `copy()` has to report a real 0 and the two would otherwise be the same number. `hash_blocks` calls `checkpoint` only on an exact position match: a forward that overshoots a rung holds state ahead of the hash it would be filed under. The interval must divide the hash block size (asserted in `BlockManager.__init__`) or a rung would have no block hash to be filed under. +**Where checkpoints land.** One ladder for every state class: a rung every `--state-checkpoint-interval-tokens` (default 8192) of context. Whether a class takes a given rung comes down to one comparison — how many tokens the *next* forward carries, against that class's `successor_room`. `BlockManager.checkpointers_at` takes the first as an argument (prefill passes what is left of the prompt, decode passes one token) and returns the classes that qualify; `checkpoint_limit` is the same rule solved for prefill's last qualifying rung and `checkpoint_cut` turns it into a chunk boundary, which the scheduler needs up front (`_finalize_prefill_chunk`). Everything else follows from the one comparison: GDN's `fork(1)` always qualifies — its `causal_conv1d` write paths all store the full window to the output slot — V4's `copy(layout_id)` reports 0 and so does too, and a rolling class needing a long hand-over simply never qualifies mid-generation. A backend with no transferable state at all declares `StateTransfer.none()`, which is `inf` on this scale; it is a separate kind rather than a token count precisely because `copy(layout_id)` has to report a real 0 and the two would otherwise be the same number. `hash_blocks` calls `checkpoint` only on an exact position match: a forward that overshoots a rung holds state ahead of the hash it would be filed under. An interval off the hash-block grid is snapped down in `BlockManager.__init__` so every rung has a block hash. **Checkpoints past the prompt.** A long answer crosses rungs the prompt never reached, and a follow-up turn replaying the conversation wants to resume from them — which is also why generated blocks enter the prefix cache at all (`hash_decode_blocks`, bounded by the committed KV length). The room test gates this with no special case: one decode token satisfies GDN's `fork(1)` and V4's `copy()` alike. Two things are gated explicitly in `Scheduler._checkpoint_room` — a request stopping on this step (nothing follows it: no forward to fork into, no batch to copy on) and speculative decode *for a forking class*. The spec exclusion has two independent reasons and either alone is decisive: the spec path's state index tensor has no read-side counterpart, so a fork must never reach it; and a spec step commits `1 + accepted_drafts` tokens, which is what a fork's successor actually gets — the rest is rolled back and re-forwarded — so no promise made when the checkpoint is decided can be kept, and by the time acceptance is known the state is already split across two groups that no single read index spans. That second reason is why DeepSeek-V4 copies rather than forks: it is the only way to checkpoint a decode boundary, which is exactly the boundary multi-turn reuse resumes from. Prefill checkpointing stays live on forking models because `min_fork_tokens` keeps prompt behind every rung and prompt always forwards down the non-spec path. Arithmetic for both compressor rings, replayed from `compress_plan.py`: `logs_claude/verify_v4_min_fork.py`. **Checkpoints where someone asked for one.** The grid is a guess about where reuse will want to resume; the requests themselves know. Whenever the state gates cut a hit short, `can_allocate` asks the same question a second time with every ladder assumed dense (`resumable_hit(..., assume_checkpointed=True)`), and the gap between the two answers is reuse that exists and is being declined only for want of a checkpoint. `BlockManager._record_checkpoint_demand` turns that into one extra rung for that seq (`Sequence.checkpoint_demand_pos`), which `checkpoint_cut` cuts a chunk at and `checkpointers_at` accepts — off the same field, so the cut and the keep cannot drift. It is decided at admission, where the counterfactual and the admitted hit are both in hand: the hit survives only as `num_cached_tokens`, which the scheduler advances as chunks land, and under pipeline parallelism is already past the chunk by the time `hash_blocks` runs. The request that discovers the gap is the one that pays for it, which is the right way round: it collects none of that reuse and has to compute the prefix anyway. The counterfactual must keep every *other* class's gate applied — a boundary some other class cannot resume from either is not worth checkpointing this one at — and demand below one interval is dropped, so a workload that keeps no checkpoints today gains no chunk cuts from this. The property is self-limiting: the first request finds nothing cached, the second finds the gap and pays one cut, and the third hits outright and finds no gap. `Lost-to-checkpoint` in the cache-stats line is the gap, measured; it falling to zero is the feature working. -**What the interval is pacing.** A checkpoint costs no capacity, but it does cost the request that takes it a forward: its prompt gets cut at the rung, and the extra forward is paid whether or not anyone ever resumes from it. That cost is the same under both mechanisms — a checkpoint holds the state as of the end of a forward, so the forward has to end on the boundary either way. What a copy adds on top is one contiguous device-to-device copy per checkpoint and per resume, which for DeepSeek-V4's ~13 MB entry is a few microseconds against a prefill measured in hundreds of milliseconds. That is why the interval counts tokens rather than blocks, and why a prompt shorter than one interval checkpoints nothing at all — on a workload of short, mutually-distinct prompts the hit rate is 0 by construction, so the feature has to be free there. Measured on Qwen3.5-27B tp2 at ISL/OSL 1024/1024, checkpointing unconditionally at the last eligible boundary cost 17.5% of total throughput for zero resumes. +**What the interval is pacing.** A checkpoint costs the request that takes it a forward: its prompt gets cut at the rung, and the extra forward is paid whether or not anyone ever resumes from it. That cost is the same under both mechanisms — a checkpoint holds the state as of the end of a forward, so the forward has to end on the boundary either way. Copy additionally consumes PAGE capacity and runs one descriptor-driven scatter or gather per checkpoint or resume; the physical PAGE ids need not be contiguous. That is why the interval counts tokens rather than blocks, and why a prompt shorter than one interval checkpoints nothing at all — on a workload of short, mutually-distinct prompts the hit rate is 0 by construction, so the feature has to be free there. Measured on Qwen3.5-27B tp2 at ISL/OSL 1024/1024, checkpointing unconditionally at the last eligible boundary cost 17.5% of total throughput for zero resumes. ### Allocation (`allocate`) @@ -345,7 +359,7 @@ def allocate(self, seq: Sequence): **Per-request cache allocation (if `seq.has_per_req_cache`):** -Pops one slot group index from the state pool's free list and assigns it to `seq.per_req_cache_group` (per-request state indexing into the builder-allocated tensors). No paged blocks are involved — the state class's bytes were already taken out of the budget at sizing time. When the hit landed on a state checkpoint, the group holding it is claimed as `seq.state_fork_src` instead and the request writes a fresh group for one forward. +Pops one slot group index from the state pool's free list and assigns it to `seq.per_req_cache_group` (per-request state indexing into the builder-allocated tensors). The resident Active Slot's bytes were already taken out of the budget at sizing time. On a fork checkpoint hit, the group holding the checkpoint is claimed as `seq.state_fork_src` and the request writes a fresh group for one forward. On a copy checkpoint hit, the request keeps its newly allocated contiguous Active Slot and queues a PAGE gather into it in the batch's `StateMaintenanceOps`. ### Deallocation (`deallocate`) diff --git a/tests/test_block_pool.py b/tests/test_block_pool.py index 188acda767..8330d63710 100644 --- a/tests/test_block_pool.py +++ b/tests/test_block_pool.py @@ -148,3 +148,47 @@ def test_a_regrown_block_is_empty_again(self): def test_growing_beyond_the_allocation_is_refused(self): with pytest.raises(ValueError, match="outside"): BlockPool(num_blocks=5, max_blocks=4) + + +class TestRawPageUnits: + def test_reservation_uses_arbitrary_non_contiguous_free_ids(self): + pool = BlockPool(num_blocks=8) + for _ in range(8): + pool.allocate(pool.pop()) + for block_id in (0, 2, 5, 7): + pool.free(block_id) + + units = pool.reserve_units(4, owner=("checkpoint", 1)) + + assert units == [0, 2, 5, 7] + assert all(pool.is_used(i) for i in units) + assert pool.num_free == 0 + with pytest.raises(AssertionError, match="belongs"): + pool.release_units(reversed(units), owner=("checkpoint", 1)) + pool.release_units(units, owner=("checkpoint", 1)) + assert pool.num_free == 4 + + def test_failed_reservation_is_atomic(self): + pool = BlockPool(num_blocks=2) + pool.allocate(pool.pop()) + assert pool.reserve_units(2, owner=("checkpoint", 1)) is None + assert pool.num_free == 1 + + def test_only_the_whole_owner_can_release_units(self): + pool = BlockPool(num_blocks=3) + units = pool.reserve_units(2, owner=("checkpoint", 1)) + with pytest.raises(AssertionError, match="belongs"): + pool.release_units(units, owner=("checkpoint", 2)) + assert pool.num_free == 1 + pool.release_units(units, owner=("checkpoint", 1)) + assert pool.num_free == 3 + + def test_retirement_refuses_a_fragment_without_relocation_protocol(self): + pool = BlockPool(num_blocks=3) + # Reserve the highest id specifically by occupying the lower two. + pool.allocate(0) + pool.allocate(1) + units = pool.reserve_units(1, owner=("checkpoint", 1)) + assert units == [2] + assert pool.retire_top() is None + assert pool.num_blocks == 3 diff --git a/tests/test_envs.py b/tests/test_envs.py index cfae983113..8b83498dca 100644 --- a/tests/test_envs.py +++ b/tests/test_envs.py @@ -28,7 +28,6 @@ "ATOM_DISABLE_VLLM_PLUGIN", "ATOM_USE_CUSTOM_ALL_GATHER", "ATOM_ENABLE_RELAXED_MTP", - "STATE_CKPT_EXTRA_ENTRIES", ] @@ -64,9 +63,6 @@ def test_dp_master_ip_default(self): def test_dp_master_port_default(self): assert _get_envs().ATOM_DP_MASTER_PORT == 29500 - def test_state_ckpt_extra_entries_default(self): - assert _get_envs().STATE_CKPT_EXTRA_ENTRIES == 0 - def test_dp_base_port_default(self): assert _get_envs().ATOM_DP_BASE_PORT == 0 @@ -116,14 +112,6 @@ def test_dp_size_override(self, monkeypatch): monkeypatch.setenv("ATOM_DP_SIZE", "8") assert _get_envs().ATOM_DP_SIZE == 8 - def test_state_ckpt_extra_entries_override(self, monkeypatch): - monkeypatch.setenv("STATE_CKPT_EXTRA_ENTRIES", "268") - assert _get_envs().STATE_CKPT_EXTRA_ENTRIES == 268 - - def test_state_ckpt_extra_entries_empty_means_default(self, monkeypatch): - monkeypatch.setenv("STATE_CKPT_EXTRA_ENTRIES", "") - assert _get_envs().STATE_CKPT_EXTRA_ENTRIES == 0 - def test_dp_port_overrides(self, monkeypatch): monkeypatch.setenv("ATOM_DP_MASTER_PORT", "29700") monkeypatch.setenv("ATOM_DP_BASE_PORT", "29800") diff --git a/tests/test_gdn_state_copy.py b/tests/test_gdn_state_relocation.py similarity index 88% rename from tests/test_gdn_state_copy.py rename to tests/test_gdn_state_relocation.py index 136aa751f8..ffa0d6638e 100644 --- a/tests/test_gdn_state_copy.py +++ b/tests/test_gdn_state_relocation.py @@ -40,11 +40,11 @@ def build(num_spec: int): @pytest.mark.parametrize("num_spec", [0, 2]) -def test_copy_moves_every_layer_and_every_slot_of_the_group(num_spec): +def test_relocation_moves_every_layer_and_every_slot_of_the_group(num_spec): stub, k, v, span = build(num_spec) before_k, before_v = k.clone(), v.clone() - GDNStateMixin.copy_state_entries(stub, [(1, 3)]) + GDNStateMixin.relocate_state_slots(stub, [(1, 3)]) src, dst = 1 * span, 3 * span assert torch.equal(k[:, dst : dst + span], before_k[:, src : src + span]) @@ -54,11 +54,11 @@ def test_copy_moves_every_layer_and_every_slot_of_the_group(num_spec): assert torch.equal(k[:, src : src + span], before_k[:, src : src + span]) -def test_copy_leaves_neighbouring_groups_alone(): +def test_relocation_leaves_neighbouring_groups_alone(): stub, k, v, span = build(num_spec=2) before_k, before_v = k.clone(), v.clone() - GDNStateMixin.copy_state_entries(stub, [(1, 3)]) + GDNStateMixin.relocate_state_slots(stub, [(1, 3)]) for group in (0, 2): lo = group * span @@ -70,7 +70,7 @@ def test_several_pairs_in_one_call(): stub, k, _, span = build(num_spec=1) before_k = k.clone() - GDNStateMixin.copy_state_entries(stub, [(0, 2), (1, 3)]) + GDNStateMixin.relocate_state_slots(stub, [(0, 2), (1, 3)]) for src, dst in ((0, 2), (1, 3)): lo_s, lo_d = src * span, dst * span @@ -81,7 +81,7 @@ def test_no_pairs_is_a_no_op(): stub, k, v, _ = build(num_spec=2) before_k, before_v = k.clone(), v.clone() - GDNStateMixin.copy_state_entries(stub, []) + GDNStateMixin.relocate_state_slots(stub, []) assert torch.equal(k, before_k) assert torch.equal(v, before_v) @@ -98,7 +98,7 @@ def test_a_group_is_not_one_slot_when_speculating(): assert span == 3 before_k = k.clone() - GDNStateMixin.copy_state_entries(stub, [(0, 2)]) + GDNStateMixin.relocate_state_slots(stub, [(0, 2)]) for offset in range(span): assert torch.equal(k[:, 2 * span + offset], before_k[:, offset]) diff --git a/tests/test_page_unit_checkpoint.py b/tests/test_page_unit_checkpoint.py new file mode 100644 index 0000000000..cbbffa9a81 --- /dev/null +++ b/tests/test_page_unit_checkpoint.py @@ -0,0 +1,206 @@ +# SPDX-License-Identifier: MIT + +"""Control-plane invariants for PAGE-backed state checkpoint images.""" + +import pickle +from dataclasses import FrozenInstanceError + +import pytest + +from atom.model_engine.block_pool import BlockPool +from atom.model_engine.page_unit_checkpoint import ( + COPYING, + EVICTING, + READY, + PagedStateCheckpointCoordinator, + PagedStateCheckpointSpec, + PageUnitCheckpointStore, +) + + +def make_store(num_units=20, unit_bytes=10, slot_bytes=25): + pool = BlockPool(num_units) + return pool, PageUnitCheckpointStore( + pool, + PagedStateCheckpointSpec( + page_unit_bytes=unit_bytes, + slot_bytes=slot_bytes, + layout_id="layout-v1", + ), + ) + + +def ready(store, prefix_hash, src_slot=0): + op = store.begin_store(prefix_hash, src_slot=src_slot) + assert op is not None + checkpoint_id = next( + cid + for cid, record in store.records.items() + if record.prefix_hash == prefix_hash + ) + assert store.records[checkpoint_id].state == COPYING + store.complete_inflight() + assert store.records[checkpoint_id].state == READY + return checkpoint_id, op + + +def test_runtime_spec_derives_units_and_has_a_minimal_wire_form(): + spec = PagedStateCheckpointSpec(10, 25, "layout-v1") + + assert spec.units_per_checkpoint == 3 + assert spec.to_wire() == { + "page_unit_bytes": 10, + "slot_bytes": 25, + "layout_id": "layout-v1", + } + assert "units_per_checkpoint" not in spec.to_wire() + assert ( + PagedStateCheckpointSpec.from_wire(pickle.loads(pickle.dumps(spec.to_wire()))) + == spec + ) + with pytest.raises(FrozenInstanceError): + spec.slot_bytes = 30 + + +@pytest.mark.parametrize( + "args", + [ + (0, 25, "layout-v1"), + (10, -1, "layout-v1"), + (10, 25, ""), + ], +) +def test_runtime_spec_rejects_invalid_geometry(args): + with pytest.raises(ValueError): + PagedStateCheckpointSpec(*args) + + +def test_runtime_spec_rejects_a_drifted_wire_shape(): + with pytest.raises(ValueError, match="fields"): + PagedStateCheckpointSpec.from_wire( + { + "page_unit_bytes": 10, + "slot_bytes": 25, + "units_per_checkpoint": 3, + "layout_id": "layout-v1", + } + ) + + +def test_copying_is_not_hash_visible_and_ready_is(): + pool, store = make_store() + op = store.begin_store(101, src_slot=3) + + assert op is not None + assert len(op.unit_ids) == 3 + assert op.total_bytes == 25 + assert store.lookup(101) == -1 + assert pool.num_free == 17 + + store.complete_inflight() + assert store.lookup(101) >= 0 + + +def test_multiple_restore_readers_pin_the_whole_record(): + pool, store = make_store() + checkpoint_id, _ = ready(store, 101) + assert store.begin_restore(101, dst_slot=4) is not None + assert store.begin_restore(101, dst_slot=8) is not None + assert store.records[checkpoint_id].pin_count == 2 + + store.unindex(101) + assert store.lookup(101) == -1 + assert store.records[checkpoint_id].state == EVICTING + assert pool.num_free == 17 + + restores = store.take_restore_ops() + assert {op.dst_slot for op in restores} == {4, 8} + store.complete_inflight() + assert checkpoint_id not in store.records + assert pool.num_free == 20 + + +def test_empty_batch_does_not_complete_a_queued_restore(): + pool = BlockPool(20) + coordinator = PagedStateCheckpointCoordinator( + pool, + PagedStateCheckpointSpec(10, 25, "layout-v1"), + enabled=True, + ) + checkpoint_id, _ = ready(coordinator.store, 101) + assert coordinator.begin_restore(101, dst_slot=4) + + coordinator.complete_previous_batch() + assert coordinator.store.records[checkpoint_id].pin_count == 1 + + _, restores = coordinator.take_checkpoint_ops() + assert len(restores) == 1 + coordinator.complete_previous_batch() + assert coordinator.store.records[checkpoint_id].pin_count == 0 + + +def test_cancel_queued_restore_drops_its_op_and_pin(): + pool, store = make_store() + checkpoint_id, _ = ready(store, 101) + assert store.begin_restore(101, dst_slot=4) is not None + store.unindex(101) + + store.cancel_queued_restore(4) + + assert store.take_restore_ops() == () + assert checkpoint_id not in store.records + assert pool.num_free == 20 + + +def test_lru_eviction_releases_one_complete_image(): + pool, store = make_store(num_units=7) + first_id, _ = ready(store, 101) + second_id, _ = ready(store, 202) + assert pool.num_free == 1 + + third = store.begin_store(303, src_slot=2) + assert third is not None + assert store.lookup(101) == -1 + assert store.lookup(202) == second_id + assert first_id not in store.records + assert store.evictions == 1 + assert len(third.unit_ids) == 3 + assert pool.num_free == 1 + + +def test_unindex_during_copy_waits_for_the_queued_writer(): + pool, store = make_store() + assert store.begin_store(101, src_slot=3) is not None + checkpoint_id = next(iter(store.records)) + store.unindex(101) + assert store.records[checkpoint_id].state == EVICTING + assert pool.num_free == 17 + + store.complete_inflight() + assert checkpoint_id not in store.records + assert pool.num_free == 20 + + +def test_protected_hit_is_excluded_from_admission_reclaim(): + pool, store = make_store(num_units=6) + ready(store, 101) + assert pool.num_free == 3 + assert store.has_available_units(6) + assert not store.has_available_units(6, protected_hash=101) + + +def test_clear_releases_ready_images_but_defers_a_pinned_reader(): + pool, store = make_store() + first_id, _ = ready(store, 101) + second_id, _ = ready(store, 202) + store.begin_restore(202, dst_slot=4) + + store.clear() + assert store.lookup(101) == store.lookup(202) == -1 + assert first_id not in store.records + assert second_id in store.records + + assert len(store.take_restore_ops()) == 1 + store.complete_inflight() + assert not store.records + assert pool.num_free == 20 diff --git a/tests/test_paged_state_copy_planner.py b/tests/test_paged_state_copy_planner.py new file mode 100644 index 0000000000..077849447e --- /dev/null +++ b/tests/test_paged_state_copy_planner.py @@ -0,0 +1,61 @@ +# SPDX-License-Identifier: MIT + +import pytest +import torch + +from atom.model_ops.attentions.paged_state_copy import ( + ByteSegment, + launch_copy_spans, + plan_segmented_copy, +) + + +def test_segmented_stream_intersection_preserves_wire_order(): + src = [ByteSegment(1000, 5), ByteSegment(2000, 7)] + dst = [ByteSegment(3000, 3), ByteSegment(4000, 4), ByteSegment(5000, 5)] + + spans = plan_segmented_copy(src, dst, total_bytes=12) + + assert [(s.src_ptr, s.dst_ptr, s.num_bytes) for s in spans] == [ + (1000, 3000, 3), + (1003, 4000, 2), + (2000, 4002, 2), + (2002, 5000, 5), + ] + + +def test_partial_tail_stops_before_unused_unit_capacity(): + src = [ByteSegment(1000, 13)] + dst = [ByteSegment(2000, 5), ByteSegment(3000, 5), ByteSegment(4000, 5)] + spans = plan_segmented_copy(src, dst, total_bytes=13) + assert sum(span.num_bytes for span in spans) == 13 + assert spans[-1].dst_ptr == 4000 + assert spans[-1].num_bytes == 3 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a GPU") +def test_descriptor_kernel_round_trips_random_bytes_with_partial_tail(): + device = torch.device("cuda") + original = torch.randint(0, 256, (13_117,), dtype=torch.uint8, device=device) + image = torch.full((14_000,), 0xA5, dtype=torch.uint8, device=device) + restored = torch.zeros_like(original) + + slot = [ByteSegment(original.data_ptr(), original.numel())] + units = [ + ByteSegment(image.data_ptr(), 4096), + ByteSegment(image.data_ptr() + 4096, 4096), + ByteSegment(image.data_ptr() + 8192, image.numel() - 8192), + ] + scatter = plan_segmented_copy(slot, units, original.numel()) + launch_copy_spans(scatter, device) + gather = plan_segmented_copy( + units, + [ByteSegment(restored.data_ptr(), restored.numel())], + original.numel(), + ) + launch_copy_spans(gather, device) + + torch.cuda.synchronize() + assert torch.equal(restored, original) + # Bytes beyond total_bytes in the final unit are never touched. + assert torch.all(image[original.numel() :] == 0xA5) diff --git a/tests/test_state_checkpoint.py b/tests/test_state_checkpoint.py index 7819ec710a..2cf28337d4 100644 --- a/tests/test_state_checkpoint.py +++ b/tests/test_state_checkpoint.py @@ -8,10 +8,9 @@ # it, a hit hands the resumed forward a group straight off the free list and it # reads the previous occupant's state. # -# Capacity model under test: a checkpoint is a FREE group whose content is -# still valid (the KV block pool's lazy eviction, applied to state groups). So -# checkpoints must never reduce the number of admissible requests, and the -# eviction event is hand-out, not free. +# Fork-transfer checkpoints are FREE groups whose content is still valid. +# Copy-transfer checkpoints are immutable PAGE-unit images; Active Slots are +# reserved only for resident requests and never serve as checkpoint backing. from math import inf, isinf from types import SimpleNamespace @@ -20,13 +19,30 @@ from conftest import MockConfig from atom.model_engine.block_manager import BlockManager +from atom.model_engine.block_pool import BlockPool +from atom.model_engine.page_unit_checkpoint import ( + PagedStateCheckpointCoordinator, + PagedStateCheckpointSpec, +) from atom.model_engine.scheduler import CacheStats, ScheduledBatchOutput, Scheduler from atom.model_engine.sequence import Sequence, SequenceType from atom.model_engine.state_cache import StateCache -from atom.model_engine.state_pool import StateGroupPool, StateTransfer +from atom.model_engine.state_pool import StateGroupPool +from atom.model_engine.state_runtime import ( + StateRuntime, + StateTransfer, +) BLOCK = 4 MIN_FORK = 8 +PAGED_COPY_SPEC = PagedStateCheckpointSpec(10, 25, "test-layout-v1") +DEFAULT_STATE_TRANSFER = StateTransfer.fork(MIN_FORK) +PAGED_COPY_TRANSFER = StateTransfer.copy(PAGED_COPY_SPEC.layout_id) +DEFAULT_STATE_RUNTIME = StateRuntime(transfer=DEFAULT_STATE_TRANSFER) +PAGED_COPY_RUNTIME = StateRuntime( + transfer=PAGED_COPY_TRANSFER, + checkpoint_spec=PAGED_COPY_SPEC, +) def ckpt_config(**overrides): @@ -43,14 +59,34 @@ def ckpt_config(**overrides): "scheduler_delay_factor": 0.0, "speculative_config": None, "pool_entries": {"state": 4}, - "state_transfer_kind": "fork", - "state_fork_tokens": MIN_FORK, "state_checkpoint_interval_tokens": BLOCK, } defaults.update(overrides) return MockConfig(**defaults) +def make_block_manager( + config, + *, + state_runtime=DEFAULT_STATE_RUNTIME, +): + return BlockManager( + config, + state_runtime=state_runtime, + ) + + +def make_scheduler( + config, + *, + state_runtime=DEFAULT_STATE_RUNTIME, +): + return Scheduler( + config, + state_runtime=state_runtime, + ) + + def stateful_seq(token_ids): return Sequence(token_ids, BLOCK, has_per_req_cache=True) @@ -84,10 +120,10 @@ def publisher_has_read_its_source(bm: BlockManager) -> None: reading and writing it at once. Tests about a resumer, not about the publisher, step over that here rather - than each spelling out two `release_state_pins` calls. + than each spelling out two lifecycle calls. """ - bm.release_state_pins() - bm.release_state_pins() + bm.complete_previous_state_batch() + bm.complete_previous_state_batch() def run_prompt_on_the_ladder(bm: BlockManager, seq: Sequence) -> list[int]: @@ -135,7 +171,7 @@ class TestPoolIndex: def test_disabled_is_identity(self): pool = StateGroupPool(0) assert pool.resumable_hit(idx_seq(), 5, [1, 2, 3, 4, 5]) == 5 - assert pool.lookup(1) == -1 + assert pool.lookup_group(1) == -1 def test_resumable_hit_picks_rightmost_checkpoint(self): pool = StateGroupPool(4, StateTransfer.fork(1), hash_block_size=1) @@ -161,20 +197,20 @@ def test_invalidate_drops_both_directions(self): pool = StateGroupPool(4) pool._index(10, 2) pool.invalidate(2) - assert pool.lookup(10) == -1 + assert pool.lookup_group(10) == -1 # A later invalidate of the same group must not delete a new tenant. pool._index(10, 3) pool.invalidate(2) - assert pool.lookup(10) == 3 + assert pool.lookup_group(10) == 3 def test_republishing_a_hash_orphans_the_old_group(self): pool = StateGroupPool(4) pool._index(10, 1) pool._index(10, 2) - assert pool.lookup(10) == 2 + assert pool.lookup_group(10) == 2 # Group 1 no longer backs hash 10; invalidating it leaves 2 indexed. pool.invalidate(1) - assert pool.lookup(10) == 2 + assert pool.lookup_group(10) == 2 def test_pins_drain_once(self): pool = StateGroupPool(4) @@ -219,7 +255,7 @@ def test_a_vacant_group_is_spent_before_any_checkpoint(self): pool.release(1) assert pool.pop() == 1 - assert pool.lookup(10) == 0 + assert pool.lookup_group(10) == 0 def test_admission_packs_towards_index_zero(self): pool = StateGroupPool(4) @@ -255,7 +291,7 @@ def test_resuming_from_a_checkpoint_refreshes_it(self): pool.release_pins() assert pool.pop() == 1 # 11 is now the older of the two - assert pool.lookup(10) == 0 + assert pool.lookup_group(10) == 0 def test_republishing_a_hash_returns_the_orphan_to_the_vacant_half(self): pool = StateGroupPool(4) @@ -266,7 +302,7 @@ def test_republishing_a_hash_returns_the_orphan_to_the_vacant_half(self): pool._index(10, 1) # group 0 no longer backs anything assert pool.pop() == 0 # vacant again, so it goes before the checkpoint - assert pool.lookup(10) == 1 + assert pool.lookup_group(10) == 1 class TestShrinking: @@ -307,8 +343,8 @@ def test_shrinking_spends_the_oldest_checkpoint_not_the_top_one(self): assert out.retired == 3 and out.held_checkpoint assert out.relocated_to == 0 - assert pool.lookup(13) == 0 # the hot one survived, at a new address - assert pool.lookup(10) == -1 # the cold one is what we spent + assert pool.lookup_group(13) == 0 # the hot one survived, at a new address + assert pool.lookup_group(10) == -1 # the cold one is what we spent assert pool.num_groups == 3 def test_the_top_is_spent_when_it_is_itself_the_oldest(self): @@ -319,7 +355,7 @@ def test_the_top_is_spent_when_it_is_itself_the_oldest(self): out = pool.retire_top() assert (out.retired, out.relocated_to, out.held_checkpoint) == (1, -1, True) - assert pool.lookup(13) == -1 + assert pool.lookup_group(13) == -1 def test_a_pinned_top_is_refused_rather_than_moved(self): """It is being read by the in-flight step; the pin drains next pass.""" @@ -370,7 +406,7 @@ def test_regrowing_a_retired_index_reuses_its_hash_slot(self): drain(pool) pool.release(2) pool._index(12, 2) - assert pool.lookup(12) == 2 + assert pool.lookup_group(12) == 2 # ── BlockManager: the hit is shrunk to a resumable boundary ──────────────── @@ -380,7 +416,7 @@ class TestHitShrink: def test_hit_is_zero_without_a_checkpoint(self): """The correctness fix: a stateful model cannot resume a bare KV hit.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) run_prompt(bm, first) # Same prompt again: compressed blocks are all cached, but the first @@ -390,10 +426,9 @@ def test_hit_is_zero_without_a_checkpoint(self): assert second.num_compressed_hit_blocks > 0 def test_stateless_model_keeps_the_full_hit(self): - bm = BlockManager( - ckpt_config( - pool_entries={}, state_transfer_kind="none", state_fork_tokens=0 - ) + bm = make_block_manager( + ckpt_config(pool_entries={}), + state_runtime=StateRuntime(), ) first = Sequence(list(range(40)), BLOCK, has_per_req_cache=False) run_prompt(bm, first) @@ -402,7 +437,7 @@ def test_stateless_model_keeps_the_full_hit(self): assert bm.can_allocate(second) == 9 def test_hit_lands_on_the_published_boundary(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) publish_at_boundary(bm, first) boundary = bm.checkpoint_limit(first) @@ -411,10 +446,10 @@ def test_hit_lands_on_the_published_boundary(self): assert bm.can_allocate(second) * bm.hash_block_size == boundary def test_resume_reads_the_checkpoint_and_writes_a_fresh_group(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) h = publish_at_boundary(bm, first) - src = bm.state.lookup(h) + src = bm.state.lookup_group(h) assert src >= 0 second = stateful_seq(list(range(40))) @@ -422,7 +457,7 @@ def test_resume_reads_the_checkpoint_and_writes_a_fresh_group(self): assert second.state_fork_src == src assert second.per_req_cache_group != src # The checkpoint survives the resume, so a third request still finds it. - assert bm.state.lookup(h) == src + assert bm.state.lookup_group(h) == src # ── Capacity: checkpoints live on the free list, never hold it back ──────── @@ -432,7 +467,7 @@ class TestCapacity: def test_checkpoints_do_not_reduce_admission(self): """A published checkpoint is a free group; concurrency is unchanged.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) for i in range(4): seq = stateful_seq(list(range(100 * i, 100 * i + 20 + 4 * i))) publish_at_boundary(bm, seq) @@ -449,10 +484,10 @@ def test_checkpoints_do_not_reduce_admission(self): assert bm.state.num_free() == 0 def test_handout_evicts_the_checkpoint_it_lands_on(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) h = publish_at_boundary(bm, first) - group = bm.state.lookup(h) + group = bm.state.lookup_group(h) bm.deallocate(first) # Drain the queue until the checkpoint's group comes back around. while bm.state.has_free(): @@ -460,16 +495,16 @@ def test_handout_evicts_the_checkpoint_it_lands_on(self): bm.allocate(seq, 0) if seq.per_req_cache_group == group: break - assert bm.state.lookup(h) == -1 + assert bm.state.lookup_group(h) == -1 def test_resume_without_a_spare_group_adopts_the_checkpoint(self): # Two groups: the publisher keeps one, so the only free group when the # resume arrives is the checkpoint itself. - bm = BlockManager(ckpt_config(pool_entries={"state": 2})) + bm = make_block_manager(ckpt_config(pool_entries={"state": 2})) first = stateful_seq(list(range(40))) h = publish_at_boundary(bm, first) publisher_has_read_its_source(bm) - group = bm.state.lookup(h) + group = bm.state.lookup_group(h) assert bm.state.num_free() == 1 second = stateful_seq(list(range(40))) @@ -478,7 +513,7 @@ def test_resume_without_a_spare_group_adopts_the_checkpoint(self): # still exactly the state it wanted, just no longer shareable. assert second.per_req_cache_group == group assert second.state_fork_src == -1 - assert bm.state.lookup(h) == -1 + assert bm.state.lookup_group(h) == -1 # ── Fork lifecycle ───────────────────────────────────────────────────────── @@ -487,7 +522,7 @@ def test_resume_without_a_spare_group_adopts_the_checkpoint(self): class TestForkLifecycle: def test_publish_moves_the_writer_to_a_new_group(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) hit = bm.can_allocate(seq) bm.allocate(seq, hit) @@ -496,10 +531,10 @@ def test_publish_moves_the_writer_to_a_new_group(self): bm.hash_blocks(seq, boundary - seq.num_cached_tokens) assert seq.per_req_cache_group != before assert seq.state_fork_src == before - assert bm.state.lookup(boundary_hash(bm, seq)) == before + assert bm.state.lookup_group(boundary_hash(bm, seq)) == before def test_no_publish_when_the_forward_misses_the_boundary(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) bm.allocate(seq, bm.can_allocate(seq)) group = seq.per_req_cache_group @@ -508,14 +543,14 @@ def test_no_publish_when_the_forward_misses_the_boundary(self): assert not bm.state.hash_to_group def test_boundary_leaves_room_for_the_fork_forward(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) boundary = bm.checkpoint_limit(seq) assert boundary % bm.hash_block_size == 0 assert seq.num_prompt_tokens - boundary >= MIN_FORK def test_every_block_boundary_up_to_the_limit_qualifies(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) limit = bm.checkpoint_limit(seq) assert bm.checkpointers_at(seq, BLOCK) @@ -526,7 +561,7 @@ def test_every_block_boundary_up_to_the_limit_qualifies(self): def test_chunked_prefill_leaves_a_ladder_of_checkpoints(self): """Intermediate boundaries publish too — the CPU-offload resume points.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) bm.allocate(seq, bm.can_allocate(seq)) for _ in range(4): @@ -534,16 +569,16 @@ def test_chunked_prefill_leaves_a_ladder_of_checkpoints(self): # the next forward, and that forward is what lets the group go. # Without the boundary four publishes would hold four sources at # once and the pool would run out mid-ladder. - bm.release_state_pins() + bm.complete_previous_state_batch() bm.hash_blocks(seq, 2 * BLOCK, start_tokens=seq.num_cached_tokens) seq.num_cached_tokens += 2 * BLOCK # Four publishes into four groups: the oldest was recycled to serve the # last one, the rest stand as distinct resume points. assert len(bm.state.hash_to_group) == 3 - assert bm.state.lookup(boundary_hash(bm, seq)) >= 0 # the rightmost one + assert bm.state.lookup_group(boundary_hash(bm, seq)) >= 0 def test_interval_thins_the_ladder(self): - bm = BlockManager(ckpt_config(state_checkpoint_interval_tokens=3 * BLOCK)) + bm = make_block_manager(ckpt_config(state_checkpoint_interval_tokens=3 * BLOCK)) seq = stateful_seq(list(range(40))) limit = bm.checkpoint_limit(seq) published = [ @@ -557,7 +592,7 @@ def test_interval_thins_the_ladder(self): assert published == [3 * BLOCK, 6 * BLOCK] def test_interval_zero_publishes_nothing(self): - bm = BlockManager(ckpt_config(state_checkpoint_interval_tokens=0)) + bm = make_block_manager(ckpt_config(state_checkpoint_interval_tokens=0)) seq = stateful_seq(list(range(40))) assert bm.checkpoint_limit(seq) == 0 assert not any(bm.checkpointers_at(seq, pos) for pos in range(BLOCK, 40, BLOCK)) @@ -569,7 +604,7 @@ def test_prompt_shorter_than_the_interval_publishes_nothing(self): request on a short-prompt workload pays an extra forward for a checkpoint nothing will ever hit. """ - bm = BlockManager(ckpt_config(state_checkpoint_interval_tokens=8 * BLOCK)) + bm = make_block_manager(ckpt_config(state_checkpoint_interval_tokens=8 * BLOCK)) seq = stateful_seq(list(range(30))) # 30 < 8 * BLOCK assert bm.checkpoint_limit(seq) == 0 run_prompt(bm, seq) @@ -586,11 +621,11 @@ def test_interval_snaps_onto_the_hash_block_grid(self): alternative the pool used to take — refusing to construct — turned a block-size choice into a startup failure naming a flag nobody set. """ - bm = BlockManager(ckpt_config(state_checkpoint_interval_tokens=BLOCK + 1)) + bm = make_block_manager(ckpt_config(state_checkpoint_interval_tokens=BLOCK + 1)) assert bm.state_checkpoint_interval_tokens == BLOCK # Below one block there is no reachable rung at all, so the ladder is # off rather than snapped to something unusable. - bm = BlockManager(ckpt_config(state_checkpoint_interval_tokens=BLOCK - 1)) + bm = make_block_manager(ckpt_config(state_checkpoint_interval_tokens=BLOCK - 1)) assert bm.state_checkpoint_interval_tokens == 0 def test_hit_never_lands_where_swa_cannot_follow(self): @@ -601,7 +636,7 @@ def test_hit_never_lands_where_swa_cannot_follow(self): somewhere SWA never approved, and `allocate` would then claim an SWA hash the pool never promised. """ - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) published = [2, 5] # checkpoint boundaries, in blocks @@ -620,13 +655,16 @@ def test_hit_never_lands_where_swa_cannot_follow(self): assert bm._gated_hit(seq, 9, hashes) == 2 def test_no_boundary_when_the_backend_cannot_fork(self): - bm = BlockManager(ckpt_config(state_transfer_kind="none", state_fork_tokens=0)) + bm = make_block_manager( + ckpt_config(), + state_runtime=StateRuntime(), + ) seq = stateful_seq(list(range(40))) assert bm.checkpoint_limit(seq) == 0 assert not bm.checkpointers_at(seq, 16) def test_cancel_adopts_the_source_and_returns_the_new_group(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) bm.allocate(seq, bm.can_allocate(seq)) source = seq.per_req_cache_group @@ -647,9 +685,9 @@ def test_two_resumers_in_one_step_share_the_checkpoint(self): # A checkpoint is read-only, so a second request hitting the same prefix # before the pins are released must fork off it too — not try to claim a # group the first one already took off the free list. - bm = BlockManager(ckpt_config(pool_entries={"state": 8})) + bm = make_block_manager(ckpt_config(pool_entries={"state": 8})) first = stateful_seq(list(range(40))) - src = bm.state.lookup(publish_at_boundary(bm, first)) + src = bm.state.lookup_group(publish_at_boundary(bm, first)) publisher_has_read_its_source(bm) resumers = [stateful_seq(list(range(40))) for _ in range(3)] @@ -664,13 +702,13 @@ def test_two_resumers_in_one_step_share_the_checkpoint(self): assert src not in groups # However many read it, the group goes back exactly once. before = bm.state.num_free() - bm.release_state_pins() + bm.complete_previous_state_batch() assert bm.state.num_free() == before + 1 def test_cancel_refuses_to_adopt_a_shared_source(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) - src = bm.state.lookup(publish_at_boundary(bm, first)) + src = bm.state.lookup_group(publish_at_boundary(bm, first)) publisher_has_read_its_source(bm) sharers = [stateful_seq(list(range(40))) for _ in range(2)] @@ -687,9 +725,9 @@ def test_cancel_refuses_to_adopt_a_shared_source(self): assert sharers[1].per_req_cache_group == src def test_cancel_of_a_resume_releases_the_pin(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) - src = bm.state.lookup(publish_at_boundary(bm, first)) + src = bm.state.lookup_group(publish_at_boundary(bm, first)) publisher_has_read_its_source(bm) second = stateful_seq(list(range(40))) @@ -699,18 +737,18 @@ def test_cancel_of_a_resume_releases_the_pin(self): assert second.per_req_cache_group == src assert not bm.state.is_pinned(src) # The pin must not also hand the group back — it has an owner now. - bm.release_state_pins() + bm.complete_previous_state_batch() assert not bm.state.is_free(src) def test_pinned_source_returns_to_the_free_list_next_step(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) - src = bm.state.lookup(publish_at_boundary(bm, first)) + src = bm.state.lookup_group(publish_at_boundary(bm, first)) publisher_has_read_its_source(bm) second = stateful_seq(list(range(40))) bm.allocate(second, bm.can_allocate(second)) assert not bm.state.is_free(src) - bm.release_state_pins() + bm.complete_previous_state_batch() assert bm.state.is_free(src) def test_a_published_source_is_not_handed_out_before_its_reader_runs(self): @@ -722,19 +760,19 @@ def test_a_published_source_is_not_handed_out_before_its_reader_runs(self): free list during the very pass that admits the requests which could pop it, and then one kernel reads and writes it at once. """ - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) - src = bm.state.lookup(publish_at_boundary(bm, first)) + src = bm.state.lookup_group(publish_at_boundary(bm, first)) assert first.state_fork_src == src assert not bm.state.is_free(src) # the pass that admits cannot get it - bm.release_state_pins() # the batch carrying the fork is built + bm.complete_previous_state_batch() # the batch carrying the fork is built assert not bm.state.is_free(src) # its forward has not been issued yet - bm.release_state_pins() # it has now + bm.complete_previous_state_batch() # it has now assert bm.state.is_free(src) # And it comes back as a checkpoint, at the LRU tail — publishing is # not what spends it. - assert bm.state.lookup(bm.state.group_hash[src]) == src + assert bm.state.lookup_group(bm.state.group_hash[src]) == src def test_a_finished_publisher_gives_its_source_back_at_once(self): """Nobody is left to read it, so the clock should not hold it. @@ -742,11 +780,11 @@ def test_a_finished_publisher_gives_its_source_back_at_once(self): This is what keeps publishing capacity-neutral for the common shape — a request that crosses a rung and then finishes or is preempted. """ - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) whole = bm.state.num_free() # nothing handed out yet h = publish_at_boundary(bm, first) - src = bm.state.lookup(h) + src = bm.state.lookup_group(h) assert not bm.state.is_free(src) bm.deallocate(first) @@ -754,7 +792,7 @@ def test_a_finished_publisher_gives_its_source_back_at_once(self): # Source and write group both back: the pool is whole again, without # waiting out the two passes the clock would have taken. assert bm.state.num_free() == whole - assert bm.state.lookup(h) == src # the checkpoint itself survives + assert bm.state.lookup_group(h) == src class TestCheckpointsDieWithTheirPrefix: @@ -768,15 +806,15 @@ class TestCheckpointsDieWithTheirPrefix: """ def test_evicting_the_block_frees_the_checkpoint_group(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) first = stateful_seq(list(range(40))) h = publish_at_boundary(bm, first) publisher_has_read_its_source(bm) - src = bm.state.lookup(h) + src = bm.state.lookup_group(h) assert bm.state.holds_checkpoint(src) bm._record_evicted(h) - assert bm.state.lookup(h) == -1 + assert bm.state.lookup_group(h) == -1 assert bm.state.is_free(src) assert not bm.state.holds_checkpoint(src) # vacant, spent before live ones assert bm.state.checkpoint_fates()["checkpoints_orphaned"] == 1 @@ -791,13 +829,13 @@ def test_an_orphan_is_spent_before_a_live_checkpoint(self): pool.unindex(10) # group 0's prefix is gone assert pool.pop() == 0 - assert pool.lookup(11) == 1 + assert pool.lookup_group(11) == 1 def test_unindex_of_an_unknown_hash_is_a_no_op(self): pool = StateGroupPool(4) pool._index(10, 0) - assert pool.unindex(999) == -1 - assert pool.lookup(10) == 0 + pool.unindex(999) + assert pool.lookup_group(10) == 0 assert pool.checkpoint_fates()["checkpoints_orphaned"] == 0 @@ -812,12 +850,12 @@ class TestPrefillChunkAlignment: """ def test_prompt_shorter_than_the_interval_is_not_cut(self): - sched = Scheduler(ckpt_config(state_checkpoint_interval_tokens=8 * BLOCK)) + sched = make_scheduler(ckpt_config(state_checkpoint_interval_tokens=8 * BLOCK)) seq = stateful_seq(list(range(30))) # 30 < 8 * BLOCK assert sched._finalize_prefill_chunk(seq, 0, 30) == 30 def test_chunk_stops_at_the_rung(self): - sched = Scheduler(ckpt_config(state_checkpoint_interval_tokens=3 * BLOCK)) + sched = make_scheduler(ckpt_config(state_checkpoint_interval_tokens=3 * BLOCK)) seq = stateful_seq(list(range(40))) limit = sched.block_manager.checkpoint_limit(seq) assert limit == 24 @@ -830,225 +868,235 @@ def test_chunk_stops_at_the_rung(self): assert sched._finalize_prefill_chunk(seq, limit, 16) == 16 -# ── Copy lifecycle ───────────────────────────────────────────────────────── +# ── PAGE-backed copy lifecycle ───────────────────────────────────────────── -def copy_config(**overrides): - """A backend whose state is one byte range: it checkpoints by copying.""" - overrides.setdefault("state_transfer_kind", "copy") - overrides.setdefault("state_fork_tokens", 0) +def paged_copy_config(**overrides): return ckpt_config(**overrides) -class TestCopyLifecycle: - """The other half of the protocol: a duplicate goes to the index. - - Everything the fork binds — a successor forward long enough to refill the - replacement, and therefore a boundary with room behind it — is gone. What - replaces it is a deferral: the bytes need a forward to move them, so the - index entry cannot appear until the copy has been scheduled. - """ - +class TestPagedCopyCheckpoint: def _admitted(self, bm, tokens=None): seq = stateful_seq(tokens or list(range(40))) bm.allocate(seq, bm.can_allocate(seq)) return seq - def test_the_owner_is_not_disturbed(self): - bm = BlockManager(copy_config()) - seq = self._admitted(bm) - group = seq.per_req_cache_group - bm.hash_blocks(seq, bm.checkpoint_limit(seq) - seq.num_cached_tokens) - # No hand-over: the group and the read slot are exactly as they were. - assert seq.per_req_cache_group == group - assert seq.state_fork_src == -1 - assert seq.pending_checkpoint != -1 - # And nothing is claimable yet — the bytes do not exist. - assert not bm.state.hash_to_group + def test_validated_runtime_is_explicit_from_wire_through_block_manager(self): + config = paged_copy_config() + engine_runtime = StateRuntime.from_wire(PAGED_COPY_RUNTIME.to_wire()) - def test_the_next_batch_turns_it_into_a_pair(self): - bm = BlockManager(copy_config()) - seq = self._admitted(bm) - src = seq.per_req_cache_group - bm.hash_blocks(seq, bm.checkpoint_limit(seq) - seq.num_cached_tokens) - h = boundary_hash(bm, seq) + scheduler = make_scheduler( + config, + state_runtime=engine_runtime, + ) - copies = bm.state_copies_for_batch() - assert seq.pending_checkpoint == -1 - assert len(copies) == 1 - got_src, dst = copies[0] - assert got_src == src and dst != src - assert bm.state.lookup(h) == dst - # Capacity-neutral: the destination went straight back on the free list. - assert bm.state.is_free(dst) - assert not bm.state_copies_for_batch() # drained once, not twice + checkpoints = scheduler.block_manager.paged_state_checkpoints + assert scheduler.block_manager.state_caches == (checkpoints,) + assert checkpoints.store.spec is engine_runtime.checkpoint_spec + assert checkpoints.store.units_per_checkpoint == 3 + assert scheduler.block_manager.state.transfer == StateTransfer.none() + assert not hasattr(scheduler.block_manager, "state_runtime") + assert not hasattr(scheduler.block_manager, "page_checkpoints") + assert not any( + hasattr(config, field) + for field in ( + "paged_state_page_unit_bytes", + "paged_state_slot_bytes", + "paged_state_units_per_checkpoint", + "paged_state_layout_id", + "state_transfer_kind", + "state_fork_tokens", + ) + ) - def test_a_request_freed_before_the_commit_indexes_nothing(self): - """Its group is back on the free list, so there is nothing to copy.""" - bm = BlockManager(copy_config()) - seq = self._admitted(bm) - bm.hash_blocks(seq, bm.checkpoint_limit(seq) - seq.num_cached_tokens) - bm.deallocate(seq) + def test_empty_batch_does_not_drain_state_maintenance(self): + scheduler = make_scheduler( + paged_copy_config(state_checkpoint_interval_tokens=0), + state_runtime=PAGED_COPY_RUNTIME, + ) + scheduler.block_manager.state.record_relocation(1, 2) - # committed by state_copies_for_batch() - assert not bm.state.hash_to_group - assert not bm.state_copies_for_batch() + scheduled = scheduler.schedule() - def test_a_full_pool_keeps_no_checkpoint(self): - """Best-effort, exactly as under a fork: no group, no checkpoint.""" - bm = BlockManager(copy_config()) - seq = self._admitted(bm) - bm.hash_blocks(seq, bm.checkpoint_limit(seq) - seq.num_cached_tokens) - while bm.state.has_free(): - bm.state.pop() + assert scheduled is None + pending = scheduler.block_manager.take_state_maintenance_ops() + assert pending.relocations == ((1, 2),) + assert scheduler.block_manager.take_state_maintenance_ops().empty - # committed by state_copies_for_batch() - assert not bm.state.hash_to_group - assert not bm.state_copies_for_batch() - - @pytest.mark.parametrize(("extra_groups", "kept"), [(0, False), (1, True)]) - def test_checkpoint_capacity_starts_above_the_live_floor(self, extra_groups, kept): - live_floor = 4 - config = copy_config( - max_num_seqs=live_floor, - pool_entries={"state": live_floor + extra_groups}, + def test_real_batch_drains_all_state_maintenance_once(self): + scheduler = make_scheduler( + paged_copy_config(state_checkpoint_interval_tokens=0), + state_runtime=PAGED_COPY_RUNTIME, ) - bm = BlockManager(config) - owners = [ - self._admitted(bm, list(range(100 * i, 100 * i + 40))) - for i in range(config.max_num_seqs) - ] - owner_groups = {seq.per_req_cache_group for seq in owners} - assert len(owner_groups) == live_floor - assert bm.state.num_free() == extra_groups - - publisher = owners[0] - bm.hash_blocks( - publisher, - bm.checkpoint_limit(publisher) - publisher.num_cached_tokens, + checkpoints = scheduler.block_manager.paged_state_checkpoints + seed = checkpoints.store.begin_store(33, src_slot=0) + assert seed is not None + checkpoints.store.complete_inflight() + assert checkpoints.begin_restore(33, dst_slot=2) + publisher = stateful_seq(list(range(BLOCK))) + publisher.per_req_cache_group = 1 + checkpoints.checkpoint(publisher, boundary_blocks=1, h=13) + scheduler.block_manager.state.record_relocation(3, 4) + scheduler.add(stateful_seq(list(range(BLOCK)))) + + batch, scheduled = scheduler.schedule() + + assert scheduled + ops = batch.state_maintenance_ops + assert ops.relocations == ((3, 4),) + assert len(ops.checkpoint_stores) == 1 + assert ops.checkpoint_stores[0].src_slot == 1 + assert len(ops.checkpoint_restores) == 1 + assert ops.checkpoint_restores[0].dst_slot == 2 + assert scheduler.block_manager.take_state_maintenance_ops().empty + assert not hasattr(scheduler.block_manager, "state_copies_for_batch") + assert not hasattr(scheduler.block_manager, "state_transfers_for_batch") + + def test_latest_pending_checkpoint_replaces_the_previous_intent(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, ) - h = boundary_hash(bm, publisher) - assert publisher.pending_checkpoint != -1 - - copies = bm.state_copies_for_batch() - assert publisher.pending_checkpoint == -1 - assert bool(copies) is kept - assert (bm.state.lookup(h) >= 0) is kept - assert bm.state.checkpoint_fates() == { - "checkpoints_kept": int(kept), - "checkpoints_dropped": int(not kept), - "checkpoints_evicted": 0, - "checkpoints_orphaned": 0, - } - if kept: - src, dst = copies[0] - assert src == publisher.per_req_cache_group - assert dst == bm.state.lookup(h) - assert dst not in owner_groups - - def test_an_existing_free_checkpoint_needs_no_second_copy(self): - config = copy_config(max_num_seqs=2, pool_entries={"state": 3}) - bm = BlockManager(config) - first = self._admitted(bm, list(range(40))) - second = self._admitted(bm, list(range(40))) + seq = self._admitted(bm) + checkpoints = bm.paged_state_checkpoints + + checkpoints.checkpoint(seq, boundary_blocks=1, h=101) + checkpoints.checkpoint(seq, boundary_blocks=2, h=202) + ops = bm.take_state_maintenance_ops() + + assert len(ops.checkpoint_stores) == 1 + assert not hasattr(seq, "pending_checkpoint") + bm.complete_previous_state_batch() + assert not checkpoints.store.contains(101) + assert checkpoints.store.contains(202) + + def test_prefix_eviction_drops_an_uncommitted_checkpoint(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, + ) + seq = self._admitted(bm) + checkpoints = bm.paged_state_checkpoints - bm.hash_blocks(first, bm.checkpoint_limit(first) - first.num_cached_tokens) - h = boundary_hash(bm, first) - copies = bm.state_copies_for_batch() - assert len(copies) == 1 - dst = copies[0][1] - assert bm.state.lookup(h) == dst - assert bm.state.is_free(dst) + checkpoints.checkpoint(seq, boundary_blocks=1, h=101) + bm._record_evicted(101) - bm.hash_blocks(second, bm.checkpoint_limit(second) - second.num_cached_tokens) - assert boundary_hash(bm, second) == h - assert bm.state_copies_for_batch() == [] - assert second.pending_checkpoint == -1 - assert bm.state.lookup(h) == dst - assert bm.state.checkpoint_fates() == { - "checkpoints_kept": 1, - "checkpoints_dropped": 0, - "checkpoints_evicted": 0, - "checkpoints_orphaned": 0, - } + assert bm.take_state_maintenance_ops().checkpoint_stores == () + assert checkpoints.checkpoint_fates()["checkpoints_orphaned"] == 1 + + def test_checkpoint_uses_page_units_not_an_active_slot(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, + ) + seq = self._admitted(bm) + free_slots = bm.state.num_free() + free_pages = bm.kv.num_free + bm.hash_blocks(seq, bm.checkpoint_limit(seq) - seq.num_cached_tokens) + h = boundary_hash(bm, seq) - def test_a_resume_is_handed_a_duplicate_not_a_fork(self): - bm = BlockManager(copy_config()) + transfers = bm.take_state_maintenance_ops() + assert transfers.relocations == () + assert len(transfers.checkpoint_stores) == 1 + assert bm.state.num_free() == free_slots + assert bm.kv.num_free == free_pages - 3 + checkpoints = bm.paged_state_checkpoints + assert checkpoints.store.lookup(h) == -1 + + bm.complete_previous_state_batch() + assert checkpoints.store.contains(h) + + def test_hit_gathers_into_a_distinct_contiguous_active_slot(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, + ) first = self._admitted(bm) bm.hash_blocks(first, bm.checkpoint_limit(first) - first.num_cached_tokens) - # committed by state_copies_for_batch() - src = bm.state_copies_for_batch()[0][1] + h = boundary_hash(bm, first) + store = bm.take_state_maintenance_ops().checkpoint_stores[0] + bm.complete_previous_state_batch() - # A follow-up turn, not a repeat: with no room reserved behind it the - # checkpoint sits on the prompt's last block, and a request of the same - # length can never reach it (its own hit stops one block short). second = stateful_seq(list(range(48))) hit = bm.can_allocate(second) assert hit > 0 bm.allocate(second, hit) - # The read side stays untouched; the bytes arrive by copy instead. + transfers = bm.take_state_maintenance_ops() + assert transfers.checkpoint_stores == () + assert len(transfers.checkpoint_restores) == 1 + restore = transfers.checkpoint_restores[0] + assert restore.unit_ids == store.unit_ids + assert restore.dst_slot == second.per_req_cache_group assert second.state_fork_src == -1 - assert bm.state_copies_for_batch() == [(src, second.per_req_cache_group)] - # And the source is held until the forward that reads it has been issued. - assert bm.state.is_pinned(src) - - def test_the_checkpoint_is_only_claimable_once_its_batch_is_decided(self): - """Why the commit waits for the batch instead of opening the pass. - - The source of a keeper copy is the owner's *live* group. Anything that - can preempt that owner between the commit and the batch — an admission, - in the same pass — would put the group back on the free list, and the - copy would then duplicate the next request's state into a group already - indexed as a checkpoint. Waiting until the batch is decided leaves no - such window, at the price of the checkpoint landing one pass later. - """ - bm = BlockManager(copy_config()) + # The checkpoint stays canonical and shareable. Its fragments were not + # adopted as the request's kernel-visible slot. + assert bm.paged_state_checkpoints.store.contains(h) + + def test_deallocate_cancels_a_queued_restore_before_reusing_its_slot(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, + ) first = self._admitted(bm) bm.hash_blocks(first, bm.checkpoint_limit(first) - first.num_cached_tokens) + h = boundary_hash(bm, first) + bm.take_state_maintenance_ops() + bm.complete_previous_state_batch() - # An admission in the same pass cannot see it yet. second = stateful_seq(list(range(48))) - assert bm.can_allocate(second) == 0 + bm.allocate(second, bm.can_allocate(second)) + dst = second.per_req_cache_group + checkpoint_id = bm.paged_state_checkpoints.store.lookup(h) + assert bm.paged_state_checkpoints.store.records[checkpoint_id].pin_count == 1 + + bm.deallocate(second) - bm.state_copies_for_batch() # the batch is decided; now it exists - assert bm.can_allocate(second) > 0 + assert bm.take_state_maintenance_ops().checkpoint_restores == () + assert bm.paged_state_checkpoints.store.records[checkpoint_id].pin_count == 0 + assert bm.state.is_free(dst) - def test_admissions_get_the_free_list_before_checkpoints_do(self): - """Committing after admissions is also the right priority order.""" - bm = BlockManager(copy_config()) + third = stateful_seq(list(range(100, 140))) + bm.allocate(third, bm.can_allocate(third)) + assert third.per_req_cache_group == dst + assert bm.take_state_maintenance_ops().checkpoint_restores == () + + def test_missing_gated_checkpoint_releases_the_new_slot_and_raises(self): + bm = make_block_manager( + paged_copy_config(), + state_runtime=PAGED_COPY_RUNTIME, + ) first = self._admitted(bm) bm.hash_blocks(first, bm.checkpoint_limit(first) - first.num_cached_tokens) - # Leave exactly one group: the admission takes it, the checkpoint yields. - while bm.state.num_free() > 1: - bm.state.pop() - - newcomer = stateful_seq(list(range(40))) - bm.allocate(newcomer, bm.can_allocate(newcomer)) - assert newcomer.per_req_cache_group >= 0 - assert bm.state_copies_for_batch() == [] - assert not bm.state.hash_to_group + h = boundary_hash(bm, first) + bm.take_state_maintenance_ops() + bm.complete_previous_state_batch() + free_slots = bm.state.num_free() + bm.paged_state_checkpoints.unindex(h) - def test_the_batch_carries_what_was_drained(self): - """The copies have to reach the forward, which means riding a batch.""" - sched = Scheduler(copy_config()) - sched.add(stateful_seq(list(range(BLOCK)))) - sched.block_manager.state.record_copy(2, 3) - batch, _ = sched.schedule() - assert batch.state_copy_pairs == [(2, 3)] - # Carried once: the next batch is not asked to repeat them. - batch, _ = sched.schedule() - assert batch.state_copy_pairs == [] + second = stateful_seq(list(range(48))) + with pytest.raises(RuntimeError, match="disappeared"): + bm._attach_state_group(second, h) - def test_a_copy_checkpoints_where_a_fork_cannot(self): - """Speculation and a one-token step both stop a fork, neither a copy.""" + assert second.per_req_cache_group == -1 + assert bm.state.num_free() == free_slots + + def test_copy_transfer_can_checkpoint_a_speculative_decode_boundary(self): spec = SimpleNamespace(num_speculative_tokens=3, use_dspark=lambda: False) seq = stateful_seq(list(range(40))) seq.type = SequenceType.DECODE - forking = Scheduler(ckpt_config(state_fork_tokens=1, speculative_config=spec)) - copying = Scheduler(copy_config(speculative_config=spec)) + forking = make_scheduler( + ckpt_config(speculative_config=spec), + state_runtime=StateRuntime(transfer=StateTransfer.fork(1)), + ) + copying = make_scheduler( + paged_copy_config(speculative_config=spec), + state_runtime=PAGED_COPY_RUNTIME, + ) + assert ( + copying.block_manager.paged_state_checkpoints.store.spec is PAGED_COPY_SPEC + ) assert forking._checkpoint_room(seq, False) == 0 assert copying._checkpoint_room(seq, False) == 1 - # A finishing request still keeps nothing: no next batch to copy on. assert copying._checkpoint_room(seq, True) == 0 @@ -1080,14 +1128,17 @@ def _prompt_of_10(self, bm): return seq def test_a_rung_past_the_prompt_publishes(self): - bm = BlockManager(ckpt_config(state_fork_tokens=1)) + bm = make_block_manager( + ckpt_config(), + state_runtime=StateRuntime(transfer=StateTransfer.fork(1)), + ) seq = self._prompt_of_10(bm) group = seq.per_req_cache_group self._generate_to(bm, seq, 3 * BLOCK) assert seq.per_req_cache_group != group assert seq.state_fork_src == group - assert bm.state.lookup(bm.kv.block(seq.block_table[2]).hash) == group + assert bm.state.lookup_group(bm.kv.block(seq.block_table[2]).hash) == group def test_a_backend_needing_a_long_fork_never_publishes_mid_generation(self): """Self-gating: no `min_fork` special case, the number decides. @@ -1095,7 +1146,7 @@ def test_a_backend_needing_a_long_fork_never_publishes_mid_generation(self): One decode token cannot fill a group that needs `MIN_FORK` of them, so the rung is simply not a publish position for this backend. """ - bm = BlockManager(ckpt_config()) # state_fork_tokens=MIN_FORK + bm = make_block_manager(ckpt_config()) # DEFAULT_STATE_TRANSFER needs MIN_FORK. seq = self._prompt_of_10(bm) group = seq.per_req_cache_group @@ -1105,7 +1156,10 @@ def test_a_backend_needing_a_long_fork_never_publishes_mid_generation(self): def test_no_publish_on_the_step_that_finishes_the_request(self): """Nothing will fork from it, and the fresh group would go straight back.""" - bm = BlockManager(ckpt_config(state_fork_tokens=1)) + bm = make_block_manager( + ckpt_config(), + state_runtime=StateRuntime(transfer=StateTransfer.fork(1)), + ) seq = self._prompt_of_10(bm) group = seq.per_req_cache_group @@ -1115,14 +1169,17 @@ def test_no_publish_on_the_step_that_finishes_the_request(self): def test_blocks_are_still_hashed_where_no_checkpoint_is_taken(self): """Prefix caching and state checkpoints are separate gates.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = self._prompt_of_10(bm) self._generate_to(bm, seq, 3 * BLOCK) assert seq.num_hashed_tokens == 3 * BLOCK def test_followup_turn_resumes_from_a_generated_rung(self): """The payoff: turn 2 reuses KV *and* the state that goes with it.""" - bm = BlockManager(ckpt_config(state_fork_tokens=1)) + bm = make_block_manager( + ckpt_config(), + state_runtime=StateRuntime(transfer=StateTransfer.fork(1)), + ) seq = self._prompt_of_10(bm) self._generate_to(bm, seq, 4 * BLOCK) @@ -1132,7 +1189,7 @@ def test_followup_turn_resumes_from_a_generated_rung(self): # left a checkpoint. assert bm.can_allocate(followup) == 3 bm.allocate(followup, 3) - assert followup.state_fork_src == bm.state.lookup( + assert followup.state_fork_src == bm.state.lookup_group( bm.kv.block(seq.block_table[2]).hash ) @@ -1141,7 +1198,10 @@ class TestDecodePublishGate: """`Scheduler._state_publish_room`: who is allowed to checkpoint at decode.""" def _sched(self, **overrides): - return Scheduler(ckpt_config(state_fork_tokens=1, **overrides)) + return make_scheduler( + ckpt_config(**overrides), + state_runtime=StateRuntime(transfer=StateTransfer.fork(1)), + ) def _decoding_seq(self): seq = stateful_seq(list(range(40))) @@ -1205,7 +1265,7 @@ def test_postprocess_carries_the_room_to_a_real_checkpoint(self): batch, _ = sched.schedule() forks.extend(s for s in batch.state_fork_srcs if s >= 0) - published = bm.state.lookup(bm.kv.block(seq.block_table[1]).hash) + published = bm.state.lookup_group(bm.kv.block(seq.block_table[1]).hash) assert published >= 0 # The seq moved off the group it gave away, and the forward right after # the publish was told to read it. @@ -1261,6 +1321,10 @@ def second_class(**overrides): class TestStateCacheProtocol: + def test_copy_transfer_has_no_slot_backed_fallback(self): + with pytest.raises(ValueError, match="do not belong"): + StateGroupPool(4, StateTransfer.copy("test-layout")) + def test_both_classes_satisfy_the_protocol(self): assert isinstance(second_class(), StateCache) assert isinstance(StateGroupPool(4), StateCache) @@ -1278,7 +1342,7 @@ def test_a_class_that_keeps_nothing_reports_inf(self): def test_the_limit_follows_the_class_that_reaches_furthest(self): """The smallest room reaches furthest right; a larger one must not cap it.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) assert bm.checkpoint_limit(seq) == 32 # the ring alone: 40 - MIN_FORK bm.state_caches = (*bm.state_caches, StubStateCache(successor_room=0)) @@ -1292,23 +1356,30 @@ def test_the_three_transfers_land_on_three_different_rooms(self): are opposite ends of the room scale. """ assert isinf(StateGroupPool(4, StateTransfer.none()).successor_room) - assert StateGroupPool(4, StateTransfer.copy()).successor_room == 0 + assert StateTransfer.copy("test-layout").successor_room == 0 assert StateGroupPool(4, StateTransfer.fork(7)).successor_room == 7 def test_a_copy_never_asks_the_resumer_for_room(self): """`resumable_hit`'s fork test is vacuous under `copy`, not skipped.""" forking = StateGroupPool(4, StateTransfer.fork(4), hash_block_size=1) - copying = StateGroupPool(4, StateTransfer.copy(), hash_block_size=1) - for pool in (forking, copying): - pool._index(10, 0) - pool._index(50, 1) + copying = PagedStateCheckpointCoordinator( + BlockPool(4), + PagedStateCheckpointSpec(1, 1, "test-layout"), + enabled=True, + ) + assert isinstance(copying, StateCache) + forking._index(10, 0) + forking._index(50, 1) + assert copying.store.begin_store(10, 0) is not None + assert copying.store.begin_store(50, 1) is not None + copying.store.complete_inflight() # Five one-token blocks; the rightmost checkpoint leaves no room to # forward, so a fork walks back to the first and a copy does not. assert forking.resumable_hit(idx_seq(5), 5, [10, 20, 30, 40, 50]) == 1 assert copying.resumable_hit(idx_seq(5), 5, [10, 20, 30, 40, 50]) == 5 def test_the_immutable_class_qualifies_where_the_rolling_one_cannot(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) # A rung one token from the end: the ring has no room to hand over, an # immutable class needs none. @@ -1319,7 +1390,7 @@ def test_the_immutable_class_qualifies_where_the_rolling_one_cannot(self): def test_cut_and_ladder_agree_position_for_position(self): """The chunk is cut where — and only where — something gets kept.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) cuts = { bm.checkpoint_cut(seq, pos - 1, pos) @@ -1337,7 +1408,7 @@ class TestGatedHitFixpoint: def test_the_answer_is_accepted_by_every_class(self): """What a fixpoint means, asserted directly rather than by construction.""" - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) hashes = [1000 + i for i in range(9)] for group, boundary in enumerate([2, 5]): @@ -1349,7 +1420,7 @@ def test_the_answer_is_accepted_by_every_class(self): assert cache.resumable_hit(seq, answer, hashes) == answer def test_order_between_classes_does_not_change_the_answer(self): - bm = BlockManager(ckpt_config()) + bm = make_block_manager(ckpt_config()) seq = stateful_seq(list(range(40))) hashes = [1000 + i for i in range(9)] for group, boundary in enumerate([2, 5]): @@ -1393,7 +1464,7 @@ class TestDemandDrivenCheckpoints: """ def test_the_gap_becomes_a_rung_off_the_grid(self): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) second = stateful_seq(PROMPT) @@ -1407,7 +1478,7 @@ def test_the_gap_becomes_a_rung_off_the_grid(self): def test_the_third_request_finds_what_the_second_was_missing(self): """Self-limiting: nothing to want, want it once, want nothing again.""" - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) first = stateful_seq(PROMPT) assert run_prompt_on_the_ladder(bm, first) == [32] # the grid alone @@ -1433,7 +1504,7 @@ def test_reuse_another_class_declines_is_not_charged_to_the_ladder(self): the whole gap to the ladder would have every request pay for a checkpoint the next one still cannot use. """ - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) bm.state_caches = (*bm.state_caches, StubStateCache(cap=8)) @@ -1454,7 +1525,7 @@ def test_a_demand_the_grid_cannot_express_is_kept_anyway(self): is no reason to discard the other. This is the workload that motivates it: prompts under the interval, sharing a real prefix. """ - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) short = list(range(16)) run_prompt_on_the_ladder(bm, stateful_seq(short)) @@ -1471,7 +1542,7 @@ def test_a_demand_the_grid_cannot_express_is_kept_anyway(self): def test_the_demand_is_cut_and_kept_at_the_same_position(self): """The cut and the keep read the same call, so they cannot drift.""" - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) seq = stateful_seq(PROMPT) bm.allocate(seq, bm.can_allocate(seq)) @@ -1490,7 +1561,7 @@ def test_a_recorded_demand_is_always_a_position_something_keeps(self): rather than argued, because the two derivations sit in different files. """ for n in range(20, 60, 3): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) tokens = list(range(1000 * n, 1000 * n + n)) run_prompt_on_the_ladder(bm, stateful_seq(tokens)) seq = stateful_seq(tokens) @@ -1499,10 +1570,9 @@ def test_a_recorded_demand_is_always_a_position_something_keeps(self): assert not demand or bm.checkpointers_at(seq, demand), n def test_a_stateless_model_records_no_demand(self): - bm = BlockManager( - demand_config( - pool_entries={}, state_transfer_kind="none", state_fork_tokens=0 - ) + bm = make_block_manager( + demand_config(pool_entries={}), + state_runtime=StateRuntime(), ) cold = Sequence(PROMPT, BLOCK, has_per_req_cache=False) run_prompt_on_the_ladder(bm, cold) @@ -1530,7 +1600,7 @@ def test_the_split_accounts_for_every_declined_token(self): def test_hit_tokens_are_counted_in_hash_blocks(self): """Under DCP one block_table entry spans `dcp` blocks of tokens.""" - sched = Scheduler(demand_config(decode_context_parallel_size=2)) + sched = make_scheduler(demand_config(decode_context_parallel_size=2)) assert sched.block_manager.hash_block_size == 2 * BLOCK seq = stateful_seq(PROMPT) seq.num_compressed_hit_blocks = 3 @@ -1560,25 +1630,25 @@ def keepers(self, bm, seq, pos, aimed): return bm.checkpointers_at(seq, pos, MIN_FORK, aimed=aimed) def test_an_aimed_step_is_held_to_the_grid(self): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) assert self.keepers(bm, seq, INTERVAL, aimed=True) assert not self.keepers(bm, seq, INTERVAL + BLOCK, aimed=True) def test_an_unaimed_step_keeps_off_the_grid(self): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) assert self.keepers(bm, seq, INTERVAL + BLOCK, aimed=False) def test_an_unaimed_step_still_has_to_land_on_a_block(self): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) # The checkpoint is filed under the hash of a whole block, so a landing # between two of them has nothing to file it under. assert not self.keepers(bm, seq, INTERVAL + 1, aimed=False) def test_spacing_is_measured_from_the_last_one_kept(self): - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) seq.last_checkpoint_pos = INTERVAL + BLOCK assert not self.keepers(bm, seq, 2 * INTERVAL, aimed=False) @@ -1587,7 +1657,7 @@ def test_spacing_is_measured_from_the_last_one_kept(self): def test_the_grid_ignores_the_watermark(self): # An aimed caller answers to `checkpoint_cut`, which knows nothing of # the watermark; letting it in here would put the two out of step. - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) seq.last_checkpoint_pos = INTERVAL assert self.keepers(bm, seq, 2 * INTERVAL, aimed=True) @@ -1597,7 +1667,7 @@ def test_a_demand_is_out_of_generation_s_reach(self): # own hit ceiling, and generation only ever asks about positions at or # past the end of the prompt. The unaimed branch omits the demand # because of this, so the day it stops holding, this fails first. - bm = BlockManager(demand_config()) + bm = make_block_manager(demand_config()) seq = stateful_seq(PROMPT) bm.allocate(seq, bm.can_allocate(seq)) second = stateful_seq(PROMPT) diff --git a/tests/test_state_transfer.py b/tests/test_state_transfer.py new file mode 100644 index 0000000000..7f228d7baa --- /dev/null +++ b/tests/test_state_transfer.py @@ -0,0 +1,124 @@ +# SPDX-License-Identifier: MIT + +"""Pure control-plane tests for backend state-transfer capabilities.""" + +import pickle + +import pytest + +from atom.model_engine.page_unit_checkpoint import PagedStateCheckpointSpec +from atom.model_engine.state_runtime import ( + StateMaintenanceOps, + StateRuntime, + StateTransfer, +) + +COPY_SPEC = PagedStateCheckpointSpec(10, 25, "layout-v2") + + +def test_wire_round_trip_keeps_the_complete_capability(): + transfers = ( + StateTransfer.none(), + StateTransfer.fork(7), + StateTransfer.copy("layout-v2"), + ) + + for transfer in transfers: + wire = pickle.loads(pickle.dumps(transfer.to_wire())) + assert StateTransfer.from_wire(wire) == transfer + assert transfers[-1].to_wire()["paged_layout_id"] == "layout-v2" + + +def test_invalid_kind_token_and_layout_combinations_are_rejected(): + invalid = ( + ("copy", 0, None), + ("copy", 1, "layout-v1"), + ("fork", 0, None), + ("fork", 1, "layout-v1"), + ("none", 1, None), + ("none", 0, "layout-v1"), + ("unknown", 0, None), + ) + + for args in invalid: + with pytest.raises(ValueError): + StateTransfer(*args) + + +def test_copy_factory_requires_a_non_empty_layout(): + with pytest.raises(ValueError, match="layout"): + StateTransfer.copy("") + + +def test_wire_shape_is_exact(): + with pytest.raises(ValueError, match="fields"): + StateTransfer.from_wire( + { + "kind": "copy", + "fork_tokens": 0, + "paged_layout_id": "layout-v1", + "other": 1, + } + ) + + +def test_state_runtime_accepts_copy_fork_and_none_contracts(): + none_runtime = StateRuntime() + fork_runtime = StateRuntime(transfer=StateTransfer.fork(7)) + copy_runtime = StateRuntime( + transfer=StateTransfer.copy(COPY_SPEC.layout_id), + checkpoint_spec=COPY_SPEC, + ) + + assert none_runtime.transfer == StateTransfer.none() + assert none_runtime.checkpoint_spec is None + assert fork_runtime.transfer.forks and fork_runtime.checkpoint_spec is None + assert copy_runtime.transfer.copies + assert copy_runtime.checkpoint_spec is COPY_SPEC + + +def test_state_runtime_rejects_every_transfer_spec_mismatch(): + invalid = ( + (StateTransfer.copy(COPY_SPEC.layout_id), None), + (StateTransfer.copy("different-layout"), COPY_SPEC), + (StateTransfer.fork(1), COPY_SPEC), + (StateTransfer.none(), COPY_SPEC), + ) + + for transfer, checkpoint_spec in invalid: + with pytest.raises(ValueError): + StateRuntime(transfer=transfer, checkpoint_spec=checkpoint_spec) + + +def test_state_runtime_wire_round_trip_revalidates_the_complete_contract(): + runtimes = ( + StateRuntime(), + StateRuntime(transfer=StateTransfer.fork(7)), + StateRuntime( + transfer=StateTransfer.copy(COPY_SPEC.layout_id), + checkpoint_spec=COPY_SPEC, + ), + ) + + for runtime in runtimes: + wire = pickle.loads(pickle.dumps(runtime.to_wire())) + assert StateRuntime.from_wire(wire) == runtime + + invalid_wire = runtimes[-1].to_wire() + invalid_wire["checkpoint_spec"] = None + with pytest.raises(ValueError, match="requires a PAGE checkpoint spec"): + StateRuntime.from_wire(invalid_wire) + + with pytest.raises(ValueError, match="fields"): + StateRuntime.from_wire({"transfer": StateTransfer.none().to_wire()}) + + +def test_state_maintenance_bundle_is_typed_and_immutable(): + empty = StateMaintenanceOps() + populated = StateMaintenanceOps(relocations=((1, 2), (3, 4))) + + assert empty.empty + assert not populated.empty + assert populated.relocations == ((1, 2), (3, 4)) + with pytest.raises(AttributeError): + populated.relocations = () diff --git a/tests/test_sub_pool_spec.py b/tests/test_sub_pool_spec.py index 7bc621f6e3..47d7e8ccd1 100644 --- a/tests/test_sub_pool_spec.py +++ b/tests/test_sub_pool_spec.py @@ -39,19 +39,7 @@ class TestStatePool: - def test_declared_extra_entries_are_preserved_without_env(self, monkeypatch): - monkeypatch.delenv("STATE_CKPT_EXTRA_ENTRIES", raising=False) - spec = state_pool(ENTRY_STATE, 10, entries_per_req=1, extra_entries=64) - assert spec.extra_entries == 64 - - @pytest.mark.parametrize(("value", "expected"), [("268", 268), ("0", 0)]) - def test_env_overrides_declared_extra_entries(self, monkeypatch, value, expected): - monkeypatch.setenv("STATE_CKPT_EXTRA_ENTRIES", value) - spec = state_pool(ENTRY_STATE, 10, entries_per_req=1, extra_entries=64) - assert spec.extra_entries == expected - - def test_empty_env_does_not_override_declared_extra_entries(self, monkeypatch): - monkeypatch.setenv("STATE_CKPT_EXTRA_ENTRIES", "") + def test_declared_extra_entries_are_preserved(self): spec = state_pool(ENTRY_STATE, 10, entries_per_req=1, extra_entries=64) assert spec.extra_entries == 64 diff --git a/tests/test_v4_sub_pool_spec.py b/tests/test_v4_sub_pool_spec.py deleted file mode 100644 index 4f5ff2887e..0000000000 --- a/tests/test_v4_sub_pool_spec.py +++ /dev/null @@ -1,88 +0,0 @@ -# SPDX-License-Identifier: MIT -# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. - -import ast -from pathlib import Path - -ROOT = Path(__file__).resolve().parent.parent -ATTENTION_DIR = ROOT / "atom" / "model_ops" / "attentions" -V4_SOURCE = ATTENTION_DIR / "deepseek_v4_attn.py" -SUB_POOL_SPEC_SOURCE = ATTENTION_DIR / "sub_pool_spec.py" -CONFIG_SOURCE = ROOT / "atom" / "config.py" -ARG_UTILS_SOURCE = ROOT / "atom" / "model_engine" / "arg_utils.py" -ENV_FIELD = "STATE_CKPT_EXTRA_ENTRIES" - - -def _runtime_field_refs(node: ast.AST, field: str) -> list[ast.AST]: - refs: list[ast.AST] = [ - child - for child in ast.walk(node) - if isinstance(child, ast.Attribute) and child.attr == field - ] - refs.extend( - child - for child in ast.walk(node) - if isinstance(child, ast.Call) - and isinstance(child.func, ast.Name) - and child.func.id == "getattr" - and len(child.args) >= 2 - and isinstance(child.args[1], ast.Constant) - and child.args[1].value == field - ) - return refs - - -def test_checkpoint_extra_entries_override_is_owned_by_state_pool(): - spec_tree = ast.parse(SUB_POOL_SPEC_SOURCE.read_text()) - state_pool_builder = next( - node - for node in spec_tree.body - if isinstance(node, ast.FunctionDef) and node.name == "state_pool" - ) - assert _runtime_field_refs(state_pool_builder, ENV_FIELD) - override_guards = [ - call - for call in ast.walk(state_pool_builder) - if isinstance(call, ast.Call) - and isinstance(call.func, ast.Attribute) - and call.func.attr == "is_set" - and call.args - and isinstance(call.args[0], ast.Constant) - and call.args[0].value == ENV_FIELD - ] - assert len(override_guards) == 1 - - tree = ast.parse(V4_SOURCE.read_text()) - builder = next( - node - for node in tree.body - if isinstance(node, ast.ClassDef) - and node.name == "DeepseekV4AttentionMetadataBuilder" - ) - method = next( - node - for node in builder.body - if isinstance(node, ast.FunctionDef) and node.name == "sub_pool_specs" - ) - state_calls = [ - call - for call in ast.walk(method) - if isinstance(call, ast.Call) - and isinstance(call.func, ast.Name) - and call.func.id == "state_pool" - and call.args - and isinstance(call.args[0], ast.Name) - and call.args[0].id == "STATE_SLOT_CLASS" - ] - assert len(state_calls) == 1 - - keywords = {kw.arg: kw.value for kw in state_calls[0].keywords} - assert ast.literal_eval(keywords["entries_per_req"]) == 1 - assert "extra_entries" not in keywords - - -def test_checkpoint_extra_entries_has_no_config_or_cli_surface(): - assert "state_checkpoint_extra_entries" not in CONFIG_SOURCE.read_text() - arg_utils = ARG_UTILS_SOURCE.read_text() - assert "state_checkpoint_extra_entries" not in arg_utils - assert "--state-checkpoint-extra-entries" not in arg_utils