diff --git a/atom/model_engine/block_manager.py b/atom/model_engine/block_manager.py index bff946aece..a01b10aedc 100644 --- a/atom/model_engine/block_manager.py +++ b/atom/model_engine/block_manager.py @@ -113,11 +113,11 @@ def __init__( checkpoint_spec = state_runtime.checkpoint_spec self.paged_state_checkpoints: PagedStateCheckpointCoordinator | None = None if checkpoint_spec is not None: + enabled = self.enable_prefix_caching and self.num_per_req_cache_groups > 0 self.paged_state_checkpoints = PagedStateCheckpointCoordinator( self.kv, checkpoint_spec, - enabled=self.enable_prefix_caching - and self.num_per_req_cache_groups > 0, + enabled=enabled, ) self.state = StateGroupPool( self.num_per_req_cache_groups, @@ -171,6 +171,7 @@ def __init__( # bug, and they are indistinguishable in the hit rate alone. self.demands_recorded: int = 0 self.chunks_cut_for_demand: int = 0 + self.demands_declined_no_room: int = 0 @classmethod def compute_hash(cls, token_ids: list[int], prefix: int = -1): @@ -212,13 +213,54 @@ def _record_evicted(self, h: int) -> None: self._state_checkpoint_cache.unindex(h) def _fresh_block(self) -> int: - """Take a block for content this step is about to compute.""" + """Take a block for content this step is about to compute. + + The raise is unreachable through `Scheduler` and the checkpoint cache + cannot make it reachable: a READY unpinned checkpoint counts as + available, and both callers sit behind a pin-aware check in the same + pass. `allocate` protects the one checkpoint it is about to pin and + sees the pins taken before it; `may_append` runs only in a pass that + scheduled no prefill (`scheduler.py`, `if num_seqs_prefill > 0` returns + first), so every pin was already released at the top of that pass. + Under contention the reachable outcome is a refused admission, not + this. + + That second half rests on prefill and decode never sharing a pass. If + the mixed batch that `scheduler.py` has a TODO for lands, `may_append` + starts running alongside this pass's pins and the argument has to be + redone. + """ 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 _checkpoint_has_room( + self, live_blocks: int = 0, protected_hash: int | None = None + ) -> bool: + """Whether an image still fits once `live_blocks` have been taken. + + `live_blocks` is what the admission asking this is about to allocate. + Counting it is the difference between "there is room for an image" and + "there is room for this request and an image", and only the second is + the question: the request's blocks are taken first. + + `protected_hash` is the checkpoint the same admission is about to pin, + excluded from what eviction could reclaim — the same argument + `can_allocate` passes to `_has_page_units` on the next line, so the + two gates in one pass agree on what is spendable. + + `True` when no PAGE-backed checkpoints exist at all: a fork checkpoint + costs the pool nothing, so there is nothing to gate. + """ + if self.paged_state_checkpoints is None: + return True + return self.paged_state_checkpoints.has_available_units( + live_blocks + self.paged_state_checkpoints.store.units_per_checkpoint, + protected_hash=protected_hash, + ) + def _has_page_units( self, count: int, protected_checkpoint_hash: int | None = None ) -> bool: @@ -363,12 +405,6 @@ def can_allocate(self, seq: Sequence) -> int: # the gates declined (compressed_hit - num_cached_blocks) from reuse # lost to compressed eviction (everything above compressed_hit). seq.num_compressed_hit_blocks = compressed_hit - self._record_checkpoint_demand( - seq, - hit=num_cached_blocks, - compressed_hit=compressed_hit, - block_hashes=block_hashes, - ) # Free-pool demand: blocks we actually reuse minus those already used # (shared ref); blocks we drop from the hit become fresh → counted. num_new_blocks = self._n_hash_blocks(seq) @@ -378,6 +414,17 @@ def can_allocate(self, seq: Sequence) -> int: protected_hash = ( block_hashes[num_cached_blocks - 1] if num_cached_blocks else None ) + # After `num_new_blocks`, not before: the demand's room check has to + # account for what this very admission is about to take, or it reads a + # pool it then drains itself. + self._record_checkpoint_demand( + seq, + hit=num_cached_blocks, + compressed_hit=compressed_hit, + block_hashes=block_hashes, + live_blocks=num_new_blocks, + protected_hash=protected_hash, + ) if not self._has_page_units(num_new_blocks, protected_hash): return -1 return num_cached_blocks @@ -634,7 +681,13 @@ def checkpoint_limit(self, seq: Sequence) -> int: return max(int((seq.num_prompt_tokens - room) // interval) * interval, 0) def _record_checkpoint_demand( - self, seq: Sequence, hit: int, compressed_hit: int, block_hashes: list[int] + self, + seq: Sequence, + hit: int, + compressed_hit: int, + block_hashes: list[int], + live_blocks: int, + protected_hash: int | None, ) -> None: """Ask the hit counterfactually, and turn the gap into a rung. @@ -678,18 +731,45 @@ def _record_checkpoint_demand( # Zero interval switches the ladder off entirely — `checkpointers_at` # keeps nothing then, so a cut for a demand would buy nothing either. interval_on = self.state_checkpoint_interval_tokens > 0 - previously_demanded = seq.checkpoint_demand_pos - seq.checkpoint_demand_pos = ( - wanted * self.hash_block_size if interval_on and wanted > hit else 0 - ) - # `can_allocate` re-runs for a sequence the queue keeps deferring, so - # count the demand when it first appears rather than once per attempt — - # otherwise one request under pressure inflates the denominator the - # convergence check above is read against. `deallocate` clears the - # field, so a re-admitted request does count again, which it should. - self.demands_recorded += bool(seq.checkpoint_demand_pos) and not ( - previously_demanded - ) + demand = wanted * self.hash_block_size if interval_on and wanted > hit else 0 + # A demand is an instruction to cut a prefill chunk onto a rung, and + # that cut costs the request a forward. Buying one for a store + # `begin_store` is about to refuse is the only part of this funnel + # that is pure loss — the attribution above stays either way, because + # the reuse really was declined for want of a checkpoint. + # + # Asked afresh on every attempt, because that is the question: a + # demand recorded while the pool had room is not still affordable once + # it does not, and letting the earlier answer stand is exactly the cut + # this gate exists to withhold. What must not repeat is the *counting*, + # which is why the seq carries its own marker rather than the gate + # reading the position it is about to overwrite. + # + # Asked with this admission's own blocks included, because they are + # taken first: a pool with room for an image but not for the request + # *and* the image would answer yes here and refuse at `begin_store`, + # with the cut already bought and the funnel showing nothing. + # + # It is still a sample. The store happens many forwards later, at the + # rung this cut creates, against a pool that has moved since — no + # question asked here can be the one `begin_store` asks. What this + # gate removes is the loss that was knowable at admission; + # `checkpoints_dropped` is what counts the rest, and the two are meant + # to be read together. + if demand and not self._checkpoint_has_room(live_blocks, protected_hash): + self.demands_declined_no_room += not seq.checkpoint_demand_declined + seq.checkpoint_demand_declined = True + demand = 0 + seq.checkpoint_demand_pos = demand + # Counted when the demand first appears rather than once per attempt — + # otherwise one deferred request inflates the denominator the + # convergence check above is read against. A separate marker from the + # decline above: a decline zeroes the position, so the position alone + # would let a recorded demand be counted twice the next time the pool + # has room. + if demand: + self.demands_recorded += not seq.checkpoint_demand_counted + seq.checkpoint_demand_counted = True def checkpoint_cut(self, seq: Sequence, start: int, end: int) -> int: """Latest ladder position in `(start, end]`, or 0 if there is none. @@ -732,6 +812,7 @@ def checkpoint_funnel(self) -> dict[str, int]: """ return { "demands_recorded": self.demands_recorded, + "demands_declined_no_room": self.demands_declined_no_room, "chunks_cut_for_demand": self.chunks_cut_for_demand, } | self._state_checkpoint_cache.checkpoint_fates() @@ -944,8 +1025,11 @@ def deallocate(self, seq: Sequence): # Covers preemption too, which frees through here and re-prefills. seq.num_hashed_tokens = 0 # Likewise the demand: it describes one admission against one cache - # state, and a re-admitted seq gets a fresh answer from `can_allocate`. + # state, and a re-admitted seq gets a fresh answer from `can_allocate` + # — including a fresh place in both funnel counters. seq.checkpoint_demand_pos = 0 + seq.checkpoint_demand_counted = False + seq.checkpoint_demand_declined = False 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. diff --git a/atom/model_engine/llm_engine.py b/atom/model_engine/llm_engine.py index fbc3d0ba03..94b08e781d 100644 --- a/atom/model_engine/llm_engine.py +++ b/atom/model_engine/llm_engine.py @@ -374,6 +374,7 @@ def get_cache_statistics(self, timeout: float = 30.0) -> dict[str, Any]: "checkpoints_dropped", "checkpoints_evicted", "demands_recorded", + "demands_declined_no_room", "chunks_cut_for_demand", ) } @@ -451,6 +452,7 @@ def summed(key: str) -> int: "checkpoints_evicted", "checkpoints_orphaned", "demands_recorded", + "demands_declined_no_room", "chunks_cut_for_demand", ) cache_totals = { diff --git a/atom/model_engine/model_runner.py b/atom/model_engine/model_runner.py index 9721d86e47..a2e2197d2a 100644 --- a/atom/model_engine/model_runner.py +++ b/atom/model_engine/model_runner.py @@ -1706,16 +1706,23 @@ def get_num_blocks(self) -> dict[str, object]: raise RuntimeError( "PAGE-backed state checkpoints require a PAGE sub-pool" ) + slot_bytes = int(plan.entry_bytes[STATE_SLOT_CLASS]) + # None means the backend has not narrowed its image: carry it all. + narrowed = self.attn_metadata_builder.checkpoint_image_bytes() checkpoint_spec = PagedStateCheckpointSpec( page_unit_bytes=int(plan.entry_bytes[plan.paged_class]), - slot_bytes=int(plan.entry_bytes[STATE_SLOT_CLASS]), + slot_bytes=slot_bytes, + image_bytes=slot_bytes if narrowed is None else int(narrowed), 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", + "slot_bytes=%d, image_bytes=%d (%.1f%% of a slot), " + "units_per_checkpoint=%d, layout=%s", checkpoint_spec.page_unit_bytes, checkpoint_spec.slot_bytes, + checkpoint_spec.image_bytes, + 100.0 * checkpoint_spec.image_bytes / checkpoint_spec.slot_bytes, checkpoint_spec.units_per_checkpoint, checkpoint_spec.layout_id, ) @@ -1865,6 +1872,10 @@ def allocate_kv_cache(self, num_kvcache_blocks): ) for name, value in per_req_state.items(): setattr(self, name, value) + # The pools are reachable through `self` only now, which is the + # earliest the builder can touch its own addresses — and the last + # moment before a request could. + self.attn_metadata_builder.warmup_per_req_cache() # Build KVCacheConfig # lirong TODO: This is a simple solution to build KVCacheConfig, diff --git a/atom/model_engine/page_unit_checkpoint.py b/atom/model_engine/page_unit_checkpoint.py index 575ddb65e6..af29852438 100644 --- a/atom/model_engine/page_unit_checkpoint.py +++ b/atom/model_engine/page_unit_checkpoint.py @@ -6,7 +6,7 @@ from __future__ import annotations from collections import OrderedDict -from collections.abc import Mapping +from collections.abc import Iterator, Mapping from dataclasses import dataclass from atom.model_engine.block_pool import BlockPool @@ -24,25 +24,40 @@ class PagedStateCheckpointSpec: page_unit_bytes: int slot_bytes: int layout_id: str + # Bytes of a slot a checkpoint image actually holds, which is less than + # all of them: a resumer reads only part of the slot it resumes into, and + # a compressor whose next pool starts exactly at the boundary reads none + # of its own. `slot_bytes` stays for three things that still want the + # whole slot — the `image_bytes <= slot_bytes` sanity check below, the + # geometry cross-check in `allocate_per_req_cache`, and the startup log + # line that reports an image as a fraction of one. + image_bytes: int def __post_init__(self) -> None: for name, value in ( ("page_unit_bytes", self.page_unit_bytes), ("slot_bytes", self.slot_bytes), + ("image_bytes", self.image_bytes), ): if not isinstance(value, int) or isinstance(value, bool) or value <= 0: raise ValueError(f"{name} must be a positive integer") + if self.image_bytes > self.slot_bytes: + raise ValueError( + f"image_bytes {self.image_bytes} exceeds the {self.slot_bytes} " + "a slot holds" + ) 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 + return (self.image_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, + "image_bytes": self.image_bytes, "layout_id": self.layout_id, } @@ -50,7 +65,7 @@ def to_wire(self) -> dict[str, int | str]: 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"} + expected = {"page_unit_bytes", "slot_bytes", "image_bytes", "layout_id"} if set(wire) != expected: raise ValueError( "invalid paged state checkpoint spec fields: " @@ -59,13 +74,14 @@ def from_wire(cls, wire: object) -> PagedStateCheckpointSpec: return cls( page_unit_bytes=wire["page_unit_bytes"], # type: ignore[arg-type] slot_bytes=wire["slot_bytes"], # type: ignore[arg-type] + image_bytes=wire["image_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.""" + """Scatter the checkpointed part of an Active Slot into PAGE units.""" src_slot: int unit_ids: tuple[int, ...] @@ -75,7 +91,7 @@ class CheckpointStoreOp: @dataclass(frozen=True) class CheckpointRestoreOp: - """Gather one ordered PAGE-unit image into an Active Slot.""" + """Gather one ordered PAGE-unit image back into an Active Slot.""" dst_slot: int unit_ids: tuple[int, ...] @@ -101,7 +117,6 @@ def __init__( ): 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] = {} @@ -134,31 +149,83 @@ def _new_identity(self) -> int: self._next_checkpoint_id += 1 return checkpoint_id + def _is_evictable(self, checkpoint_id: int, protected: int = -1) -> bool: + """Whether this checkpoint may be spent. Eligibility, not policy. + + The one statement of what is evictable. Everything that asks about + free units goes through it, so a new state, a grace period or a + second kind of pin cannot leave two answers behind. Do not order + here -- which eligible checkpoint to spend first is `_next_victim`. + """ + record = self.records[checkpoint_id] + return ( + checkpoint_id != protected + and record.state == READY + and record.pin_count == 0 + ) + + def _evictable(self, protected: int = -1) -> Iterator[int]: + """Every checkpoint that may be spent. Yield order carries no promise.""" + return (cid for cid in self._lru if self._is_evictable(cid, protected)) + + def _next_victim(self, protected: int = -1) -> int: + """Which eligible checkpoint to spend when the free list is short. + + This is the eviction policy, and the only place it lives: least + recently used, which `_lru` already orders. A different policy + replaces this method and nothing else -- in particular it must not + touch `_is_evictable`, which is the eligibility rule three callers + share. + """ + return next(self._evictable(protected), -1) + def has_available_units( self, count: int, protected_hash: int | None = None ) -> bool: + """Whether `count` units could be had, evicting if it came to that. + + Asked once per waiting sequence in `can_allocate` and once per running + one in `can_append`, so it is per-sequence per-pass and the walk has to + be paid for. Two things keep it cheap. The free list is checked first, + which is the whole answer whenever the pool is not tight. And the walk + below stops at the shortfall rather than totalling the cache: the + question is whether the eligible set reaches `count`, not how large it + is, and a warm pool holds `num_kvcache_blocks / units_per_checkpoint` + checkpoints -- thousands, walked for an answer a couple of them settle. + + Which checkpoints those are does not change the answer, only how soon + the loop reaches it, so a future `_next_victim` cannot move this gate. + """ 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 + shortfall = count - self.pool.num_free + for checkpoint_id in self._evictable(protected): + shortfall -= len(self.records[checkpoint_id].unit_ids) + if shortfall <= 0: + return True + return False def ensure_free_units(self, count: int) -> bool: + """Raise the free list to `count`, spending checkpoints for the shortfall. + + Free units are taken first -- a caller asking for what is already + there evicts nothing -- and `pop` hands out never-used blocks before + cached ones, so a store reaches for the cache only once the pool has + nothing spare. Each eviction returns a whole image's units, so the + loop overshoots by at most one checkpoint. + + Unreachable counts are refused before anything is spent. The loop + alone gives up only once it has evicted everything it can, so a count + the cache cannot reach would destroy the cache on the way to saying + no. The test lives here rather than in the one caller that used to + carry it, because every caller needs it and only the argument being + 1 keeps `_fresh_block` from needing it today. + """ + if not self.has_available_units(count): + return False 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, - ) + victim = self._next_victim() if victim < 0: return False self._evict(victim) @@ -168,6 +235,23 @@ def begin_store(self, prefix_hash: int, src_slot: int) -> CheckpointStoreOp | No if self.lookup(prefix_hash) >= 0 or prefix_hash in self._pending_by_hash: return None needed = self.units_per_checkpoint + # A store takes what its own image needs and nothing more. It used to + # take a floor for live KV on top, which meant one accepted store + # spent tens of checkpoints to build a cushion -- and the cushion + # bought nothing: the pool cannot starve live KV. A READY unpinned + # checkpoint is already counted as available by `has_available_units`, + # so holding one costs live KV nothing; the unevictable set (COPYING, + # or pinned by a restore) is created after every allocation in a pass + # and resolved before the next one allocates; and every `_fresh_block` + # sits behind a pin-aware check in its own pass, so the reachable + # outcome is a refused admission, never the raise. + # + # A store that will be dropped has to cost nothing, which is what + # `ensure_free_units` refusing before it evicts buys. Its answer is + # read rather than assumed: `_next_victim` is meant to be replaced, + # and a policy that passes over an eligible checkpoint would leave the + # loop short after spending some -- taking an identity and a record + # for a store that cannot happen would then be the second cost. if not self.ensure_free_units(needed): return None @@ -186,7 +270,7 @@ def begin_store(self, prefix_hash: int, src_slot: int) -> CheckpointStoreOp | No return CheckpointStoreOp( src_slot=src_slot, unit_ids=record.unit_ids, - total_bytes=self.spec.slot_bytes, + total_bytes=self.spec.image_bytes, layout_id=self.spec.layout_id, ) @@ -202,7 +286,7 @@ def begin_restore( op = CheckpointRestoreOp( dst_slot=dst_slot, unit_ids=record.unit_ids, - total_bytes=self.spec.slot_bytes, + total_bytes=self.spec.image_bytes, layout_id=self.spec.layout_id, ) self._queued_restores.append((checkpoint_id, op)) diff --git a/atom/model_engine/sequence.py b/atom/model_engine/sequence.py index ca139f97e3..8bce6c5d77 100644 --- a/atom/model_engine/sequence.py +++ b/atom/model_engine/sequence.py @@ -105,6 +105,15 @@ def __init__( # `BlockManager._record_checkpoint_demand` at admission; this one is # read by `checkpoint_cut` and `checkpointers_at`, which must agree. self.checkpoint_demand_pos = 0 + # Which of the two demand counters this seq has already been put + # against. `can_allocate` re-runs for a sequence the queue keeps + # deferring, and the position above cannot serve as the marker for + # either one: a declined demand writes 0 back, so it does not remember + # the decline, and a decline retracts a demand the recorded counter had + # already taken. Both are cleared by `deallocate`, so a re-admitted + # request counts again, which it should. + self.checkpoint_demand_counted = False + self.checkpoint_demand_declined = False # Where this seq last kept a checkpoint. Prefill lands on the grid so # this tracks it, but a speculative decode step lands wherever # `1 + accepted` puts it, and there the grid is unreachable — see diff --git a/atom/model_ops/attentions/backends.py b/atom/model_ops/attentions/backends.py index 2ab2ddfa68..adba545ae4 100644 --- a/atom/model_ops/attentions/backends.py +++ b/atom/model_ops/attentions/backends.py @@ -147,6 +147,17 @@ def state_transfer(self) -> StateTransfer: """Declare this backend's per-request state checkpoint capability.""" return StateTransfer.none() + def checkpoint_image_bytes(self) -> int | None: + """Bytes of an Active Slot a checkpoint image has to hold. + + `None` means all of them: the safe answer, and the one a backend that + has not worked out which of its bytes a resumer skips should keep + giving. A backend returns less only when it can name bytes no resumer + reads — for a ring whose next reader starts exactly at the checkpoint + boundary, that is the whole ring. + """ + return None + def relocate_state_slots(self, pairs: Sequence[tuple[int, int]]) -> None: """Move live state between contiguous Active Slots.""" raise NotImplementedError( @@ -165,6 +176,16 @@ def execute_paged_state_copies( f"{type(self).__name__} does not implement PAGE-backed state copy" ) + def warmup_per_req_cache(self) -> None: + """Pay whatever the first checkpoint copy would pay, before serving. + + Called once by ModelRunner after `allocate_per_req_cache`'s pools are + installed, which is the earliest a backend can reach its own addresses. + Nothing else warms this path: `execute_paged_state_copies` runs only + from `build()`, so a backend that compiles a kernel or fills a cache + there does it inside a live request's batch. A no-op by default. + """ + def get_kv_transfer_tensors(self) -> "KVTransferTensors | None": """Return RDMA transfer regions for PD disaggregation. diff --git a/atom/model_ops/attentions/deepseek_v4_attn.py b/atom/model_ops/attentions/deepseek_v4_attn.py index 7e95291844..e92b2a09de 100644 --- a/atom/model_ops/attentions/deepseek_v4_attn.py +++ b/atom/model_ops/attentions/deepseek_v4_attn.py @@ -70,10 +70,16 @@ AttentionMetadataBuilder, CommonAttentionBuilder, ) +from atom.model_ops.attentions.paged_state_copy import ( + SegmentedCopyPlan, + launch_copy_descriptor, + plan_segmented_copy, +) from atom.model_ops.attentions.state_arena import ( SplitStateArena, StateArena, StateField, + checkpoint_ranges_for, plan_field_planes, plan_regions, ) @@ -89,6 +95,7 @@ HCA_RATIO, UnifiedPoolGeometry, WindowParams, + merge_abutting, ) from atom.model_ops.v4_kernels import ( FP4_MQA_BLOCK_K, @@ -606,8 +613,18 @@ def __init__(self, model_runner): self._alloc_v4_metadata_buffers() self._ubatch_decode_meta: list | None = None - # Filled on the first checkpoint copy — the pools do not exist yet here. + # Filled on the first checkpoint copy — the pools do not exist yet + # here. Four of the five hold raw addresses read out of the pool + # tensors, so a re-carve invalidates them all; that is what + # `_invalidate_pool_caches` is for, and `allocate_per_req_cache` + # calls it. self._slot_view_cache: list[list[torch.Tensor]] | None = None + self._checkpoint_range_cache: list[list[tuple[int, int]]] | None = None + self._page_unit_region_cache: tuple[np.ndarray, np.ndarray] | None = None + self._page_unit_region_owners: tuple[int, ...] = () + self._checkpoint_plan_cache: SegmentedCopyPlan | None = None + self._checkpoint_slot_base_cache: np.ndarray | None = None + self._checkpoint_descriptor: CpuGpuBuffer | None = None @property def prep_stream(self): @@ -748,8 +765,28 @@ def _state_fields(self) -> list[StateField]: StateField("csa_main_score", n_csa, self.csa_main_state_shape, dt, neg_inf), StateField("csa_idx_kv", n_csa, self.csa_idx_state_shape, dt), StateField("csa_idx_score", n_csa, self.csa_idx_state_shape, dt, neg_inf), - StateField("hca_main_kv", n_hca, self.hca_main_state_shape, dt), - StateField("hca_main_score", n_hca, self.hca_main_state_shape, dt, neg_inf), + # HCA owes a checkpoint nothing. It pools `ratio` tokens with no + # overlap, so the first compression at or after a boundary P + # covers `[P, P + 128)` — every row of it written by the very + # forward that reads it — and a checkpoint sits on a multiple of + # `hash_block_size`, which `_assert_ratios_divide_the_alignment` keeps a + # multiple of 128. The rows past `K_pool` are speculative + # rollback slack and are never read at all. + StateField( + "hca_main_kv", + n_hca, + self.hca_main_state_shape, + dt, + in_checkpoint=False, + ), + StateField( + "hca_main_score", + n_hca, + self.hca_main_state_shape, + dt, + neg_inf, + in_checkpoint=False, + ), ] if self._field_window_dtype is not None: fields.append( @@ -769,8 +806,15 @@ def _state_fields(self) -> list[StateField]: def state_transfer(self) -> StateTransfer: """Declare PAGE-copy checkpoints with the versioned DSV4 layout.""" ratios = ",".join(str(r) for r in self._geometry_ratios()) + nocopy = ",".join(f.name for f in self._state_fields() if not f.in_checkpoint) layout_id = ( - "dsv4-paged-state-v1" + # v2: an image holds part of a slot, not all of it. Two workers + # disagreeing about which part would read one image at two + # layouts, so `nocopy` names the rule and the version fences it. + # v3: it also drops the entry's interleave padding, so the image + # is no longer a subsequence of the slot's rows and a v2 reader + # would gather every window row shifted. + "dsv4-paged-state-v3" 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}," @@ -778,6 +822,8 @@ def state_transfer(self) -> StateTransfer: f":main={'fp8-2buff' if self._kv_fp8 else 'bf16'}" f":index={'fp4' if self._indexer_fp4 else 'fp8'}" f":ratios={ratios}" + f":nocopy={nocopy}" + ":entry=packed" ) return StateTransfer.copy(layout_id) @@ -796,25 +842,45 @@ def execute_paged_state_copies( 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, - ) + """Copy raw checkpoint bytes between slots and non-contiguous PAGEs. - 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: + Every op of either direction goes into one descriptor and one launch. + A store and a restore are the same intersection read opposite ways, so + they share the plan too — and each direction is described in a single + vectorised pass, which is why they are batched apart rather than + interleaved. + """ + if not store_ops and not restore_ops: + return + for op in (*store_ops, *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) + + plan = self._checkpoint_copy_plan() + slot_bases = self._checkpoint_slot_bases() + per_op = plan.num_spans + total = (len(store_ops) + len(restore_ops)) * per_op + staging = self._checkpoint_descriptor_buffer() + if total > staging.np.shape[0]: + raise RuntimeError( + f"a step asked to copy {total // per_op} checkpoints, more " + f"than the {staging.np.shape[0] // per_op} its descriptor was " + "sized for" + ) + descriptor = staging.np[:total] + at = 0 + for ops, storing in ((store_ops, True), (restore_ops, False)): + if not ops: + continue + end = at + len(ops) * per_op + groups = [op.src_slot if storing else op.dst_slot for op in ops] + plan.write_descriptor( + descriptor[at:end], + slot_bases[groups], + self._page_unit_bases([op.unit_ids for op in ops]), + forward=storing, + ) + at = end + launch_copy_descriptor(staging.copy_to_gpu(total), plan) def _validate_paged_state_op( self, op: CheckpointStoreOp | CheckpointRestoreOp @@ -827,10 +893,10 @@ def _validate_paged_state_op( f"state checkpoint layout mismatch: {op.layout_id!r} != " f"{spec.layout_id!r}" ) - if op.total_bytes != spec.slot_bytes: + if op.total_bytes != spec.image_bytes: raise RuntimeError( f"state checkpoint size mismatch: op={op.total_bytes}, " - f"active_slot={spec.slot_bytes}" + f"image={spec.image_bytes}" ) if len(op.unit_ids) != spec.units_per_checkpoint: raise RuntimeError( @@ -840,38 +906,316 @@ def _validate_paged_state_op( 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 + def _checkpoint_slot_ranges(self) -> list[list[tuple[int, int]]]: + """Per plane, the `(offset, num_bytes)` of a slot a checkpoint holds. - return [tensor_segment(view) for view in self._slot_views()[group]] + A slot is a request's compressor state and then its sliding windows + (`v4_pool_geometry`). The two halves answer differently: a window is a + sliding window, so a resumer needs every row of it, while most of the + state is dead at a boundary and says so through + `StateField.in_checkpoint`. - 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 + Three kinds of byte belong to neither and are left out: the padding + the state's byte count is rounded up by, the slot's own tail + alignment, and the entry's interleave padding — rows no `ring_row` + reaches, so no window is missing anything without them + (`UnifiedPoolGeometry.entry_row_runs`). - 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)) + A property of the layout, not of any one slot, so it is computed once. + """ + if self._checkpoint_range_cache is None: + self._assert_ratios_divide_the_alignment() + geo = self.pool_geometry + # Rows, so the same for every plane; only the width they are + # priced at differs. + window_runs = geo.entry_row_runs() + self._checkpoint_range_cache = [ + # The last state field can end exactly where the entry begins. + merge_abutting( + [ + *checkpoint_ranges_for(fields), + *( + ((geo.arena_rows + start) * width, count * width) + for start, count in window_runs + ), + ] + ) + for fields, width in zip( + self._arena_planes, self._plane_row_widths(), strict=True + ) + ] + return self._checkpoint_range_cache + + def _checkpoint_segment_sizes(self) -> list[int]: + """The image as the copy planner reads it: one size per source segment. + + The one place the per-plane ranges are flattened. Sizing wants their + total and the planner wants the list, and the two answering from + different comprehensions is how an image gets priced at one shape and + cut at another. + """ + return [ + nbytes for ranges in self._checkpoint_slot_ranges() for _, nbytes in ranges + ] + + def checkpoint_image_bytes(self) -> int: + """Bytes one checkpoint image holds. Priced before the pool exists.""" + return sum(self._checkpoint_segment_sizes()) + + def _assert_ratios_divide_the_alignment(self) -> None: + """A checkpoint boundary has to be a compression boundary too. + + `StateField.in_checkpoint` says a compressor without overlap owes a + checkpoint nothing, because the first pool at or after the boundary + starts exactly on it. That holds only while every ratio divides the + quantity a checkpoint is aligned to, and HCA's ratio is 128 — let the + alignment drop under that and HCA silently starts needing rows it is + no longer given. No crash, just a resumer reading stale KV for its + first pool, which reads as a small accuracy loss and nothing else. + + The quantity is `BlockManager`'s `hash_block_size`, not this class's + own `block_size`: the ladder rounds a checkpoint to the prefix-cache + hash granularity, which is `kv_cache_block_size * dcp_world_size`. + + The ratios come from the model rather than from the two constants this + file happens to name, and that is what makes the guard reachable at + all. `config.py` forces `kv_cache_block_size` to 256 for every + `DeepseekV4*` architecture, so the alignment is always a multiple of + 4 and of 128 and a check written against `CSA_RATIO`/`HCA_RATIO` could + never fire — it would only restate what the config already pins. What + is genuinely free to change is `hf_config.compress_ratios`: a variant + that pools on some other stride is the edit that breaks the premise, + and this is what catches it. + """ + config = self.model_runner.config + # Raw, like `BlockManager.__init__` reads it: a zero would give an + # alignment of zero, which every ratio divides, so a `or 1` here would + # answer a question about a pool geometry that cannot exist. + alignment = int(config.kv_cache_block_size) * int( + config.decode_context_parallel_size ) + # Its own check, not folded into the one below: every ratio divides + # zero, so a non-positive alignment reaches that test with nothing to + # report and would raise saying it is not a multiple of `[]`. + if alignment <= 0: + raise ValueError( + f"a checkpoint aligns to {alignment} tokens " + "(kv_cache_block_size x decode_context_parallel_size), which " + "is not a pool geometry that can exist" + ) + bad = sorted({r for r in self.compress_ratios if r > 0 and alignment % r}) + if bad: + raise ValueError( + f"a checkpoint aligns to {alignment} tokens " + f"(kv_cache_block_size x decode_context_parallel_size), which " + f"is not a multiple of compress ratios {bad}, so a checkpoint " + "boundary is not a compression boundary, so what " + "`_state_fields` leaves out of the image is no longer dead" + ) + + def _invalidate_pool_caches(self) -> None: + """Forget everything derived from the pools' layout or addresses. + + A slot's address is a function of the *split*, not just of its group + (`UnifiedPoolGeometry.physical_slot` counts back from the topmost + position), so a re-carve moves every one of them. Whoever wires an + elastic pool has to call this; it is here so that is one line rather + than a list of fields to remember. + + `_page_unit_region_cache` is deliberately not in it: half of what it + holds comes from pools this method's caller does not own, so being on + the list would make it look covered when it is not. It keys on its own + addresses instead. + """ + self._slot_view_cache = None + self._checkpoint_range_cache = None + self._checkpoint_plan_cache = None + self._checkpoint_slot_base_cache = None + self._checkpoint_descriptor = None + + def warmup_per_req_cache(self) -> None: + """Run one checkpoint copy now, so the first real one is only a copy. + + `execute_paged_state_copies` is reachable only from `build()`, so + everything it builds lazily -- the copy plan, the slot views, the slot + base table, the tiling's upload, the pinned descriptor, and the Triton + JIT of `_copy_tiles_kernel` -- otherwise lands inside the batch of + whichever request first crosses a rung. Hundreds of milliseconds, once, + on one unlucky request. + + Slot 0 into the pool's first units. Both are real addresses, which is + the point: a warmup on scratch would compile a kernel and fill nothing. + The bytes it writes are read by nobody -- a KV block is written before + it is read, and this runs before any block has been handed out. + """ + if self.model_runner.state_runtime.checkpoint_spec is None: + return + plan = self._checkpoint_copy_plan() + if not plan.num_spans: + return + units = self.model_runner.state_runtime.checkpoint_spec.units_per_checkpoint + staging = self._checkpoint_descriptor_buffer() + plan.write_descriptor( + staging.np[: plan.num_spans], + self._checkpoint_slot_bases()[:1], + self._page_unit_bases([list(range(units))]), + ) + launch_copy_descriptor(staging.copy_to_gpu(plan.num_spans), plan) + + def _checkpoint_descriptor_buffer(self) -> CpuGpuBuffer: + """Pinned staging for a step's whole descriptor, sized for the worst step. + + Pinned because the alternative synchronizes: a pageable H2D from + `build()` makes the host wait out the forward already enqueued, which + measured 2.9 ms behind 4 ms of work against 0.1 ms staged. Reused + because allocating pinned memory is itself a synchronizing call. + + A step can carry at most one store and one restore per sequence, so + two per Active Slot bounds it -- 1.6 MB at the shipped geometry. The + caller checks that bound rather than growing on demand: a descriptor + that did not fit would otherwise be silently truncated into a copy of + the wrong shape. + """ + if self._checkpoint_descriptor is None: + plan = self._checkpoint_copy_plan() + max_ops = 2 * int(self.model_runner.config.max_num_seqs) + self._checkpoint_descriptor = CpuGpuBuffer( + max_ops * plan.num_spans, + 3, + dtype=torch.int64, + device=self._kv_planes()[0].device, + ) + return self._checkpoint_descriptor + + def _checkpoint_copy_plan(self) -> SegmentedCopyPlan: + """Where a slot's checkpoint ranges meet a whole image's PAGE regions. + + Both streams are geometry. The ranges come from the layout, and every + image is `units_per_checkpoint` units of identical region sizes — + `_validate_paged_state_op` refuses anything else. So the cut points are + the same for every store and every restore this worker will ever do, + and the walk that finds them runs once instead of once an op. + """ + if self._checkpoint_plan_cache is None: + spec = self.model_runner.state_runtime.checkpoint_spec + self._checkpoint_plan_cache = plan_segmented_copy( + self._checkpoint_segment_sizes(), + # Sizes from the same array `_page_unit_bases` takes addresses + # from, tiled the way it ravels. Spelling the destination + # stream out a second time here would let the two orders + # diverge, and a plan cut against one order and addressed + # through the other lands whole regions in the wrong unit. + self._page_unit_stream_sizes(spec.units_per_checkpoint), + spec.image_bytes, + ) + return self._checkpoint_plan_cache + + def _checkpoint_slot_bases(self) -> np.ndarray: + """`[group, segment]` start address of every source segment of a copy. + + One row per pool group, segments in the order + `_checkpoint_slot_ranges` walks the planes — which is the order + `_checkpoint_copy_plan` built the source stream in, so a plan's + segment indices address a row of this directly. + + Materialized rather than recomputed because a group's slot sits at a + fixed address for the pool's whole life, which leaves the entire + per-op source side as one row lookup. + """ + if self._checkpoint_slot_base_cache is None: + self._checkpoint_slot_base_cache = np.array( + [ + [ + view.data_ptr() + offset + for view, ranges in zip( + views, self._checkpoint_slot_ranges(), strict=True + ) + for offset, _ in ranges + ] + for views in self._slot_views() + ], + dtype=np.int64, + ) + return self._checkpoint_slot_base_cache + + def _page_unit_regions(self) -> tuple[np.ndarray, np.ndarray]: + """Base address and per-block stride of every region a PAGE id owns. + + Blocks sit back to back in every pool, so a block's stride there is + its size and its address is `base + block_id * num_bytes` — affine, + and a property of the pools rather than of any block, so it is worked + out once. Slicing the tensors instead, which is what this replaced, + built 22 views per unit per op and threw them away to learn addresses + that are one multiplication each. + + Two arrays rather than a list of pairs because both callers want + columns: one multiplies the strides by an id, the other tiles them. + + The contiguity a slice would have been checked for is asked here + instead: once of the layout, rather than every time of a slice. + + Keyed on the addresses it was built from rather than cleared by + `_invalidate_pool_caches`. Half of these come from the KV planes, + which that hook covers, and half from the indexer pools, which + `allocate_kv_cache_tensors` owns and it does not -- so the hook would + be an invariant this cache cannot check and the next reader cannot + see. The key is the two planes and the one or two indexer pools, so + finding out costs four `data_ptr()`s against a copy path measured in + milliseconds, and the + failure it removes is a scatter into whatever the allocator handed + that address range to next. + """ + runner = self.model_runner + planes = self._kv_planes() + pools = [runner.v4_csa_idx_kv] 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)) + pools.append(runner.v4_csa_idx_kv_scale) + owners = tuple(t.data_ptr() for t in (*planes, *pools)) + # The whole test: there is always at least one plane, so the owners of + # a built cache are never the empty tuple this starts as, and an + # `is None` beside this would be a second condition that cannot differ + # from it. + if self._page_unit_region_owners != owners: + geo = self.pool_geometry + bases, strides = [], [] + for plane, width in zip(planes, self._plane_row_widths(), strict=True): + if not plane.is_contiguous(): + raise RuntimeError("a KV plane must be contiguous to be copied") + bases.append(plane.data_ptr()) + strides.append(geo.envelope_rows * width) + for pool in pools: + if not pool.is_contiguous(): + raise RuntimeError( + "an indexer pool must be contiguous to be copied" + ) + per_layer = pool.stride(0) * pool.element_size() + per_block = pool.stride(1) * pool.element_size() + for layer in range(len(self.csa_layers)): + bases.append(pool.data_ptr() + layer * per_layer) + strides.append(per_block) + self._page_unit_region_cache = ( + np.array(bases, dtype=np.int64), + np.array(strides, dtype=np.int64), ) - return segments + self._page_unit_region_owners = owners + return self._page_unit_region_cache - 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 _page_unit_bases(self, unit_ids: Sequence[Sequence[int]]) -> np.ndarray: + """Start address of every destination segment, one row per image. + + `unit_ids` is `(images, units_per_checkpoint)`. A unit's regions are + each at `base + id * stride`, so one image's worth is an outer product + and a batch's is the same product with an image axis in front. Unit + major, region minor — the order `_checkpoint_copy_plan` built the + destination stream in. + """ + base, stride = self._page_unit_regions() + ids = np.asarray(unit_ids, dtype=np.int64) + return (base + ids[..., None] * stride).reshape(len(ids), -1) + + def _page_unit_stream_sizes(self, units: int) -> np.ndarray: + """Bytes in each destination segment of an image of `units` units.""" + return np.tile(self._page_unit_regions()[1], units) def _slot_views(self) -> list[list[torch.Tensor]]: """Per-group views of that request's whole slot in each plane. @@ -1077,6 +1421,18 @@ def allocate_per_req_cache(self, entries: dict[str, int]) -> dict[str, object]: dtype = self._swa_dtype rope_dtype = self._rope_dtype + # Anything already worked out from the old layout or the old pools is + # now wrong, and wrong quietly: four of these hold raw addresses, so a + # stale one is a copy to the wrong slot rather than a crash. Cleared + # here, before `pool_geometry` is replaced, so that nothing between + # this line and that one can answer from the old split. + # + # Below that line they refill freely, and the cross-check further down + # depends on it: `checkpoint_image_bytes()` is what re-derives the + # ranges, and it has to derive them from the geometry just installed. + # What is not allowed is a *second* clear after that point — it would + # throw away the answer the check just validated. + self._invalidate_pool_caches() # The layout at the split sizing chose. Everything below — and every # index formula any kernel evaluates — reads its offsets from here. geo = self.pool_geometry.with_capacity(num_blocks, num_slots) @@ -1094,15 +1450,18 @@ def allocate_per_req_cache(self, entries: dict[str, int]) -> dict[str, object]: if checkpoint_spec is None: raise RuntimeError("DSV4 PAGE/state checkpoint sizing spec is missing") layout_id = state_runtime.transfer.paged_layout_id + actual_image_bytes = self.checkpoint_image_bytes() if ( actual_page_bytes != checkpoint_spec.page_unit_bytes or actual_slot_bytes != checkpoint_spec.slot_bytes + or actual_image_bytes != checkpoint_spec.image_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"image={actual_image_bytes}/{checkpoint_spec.image_bytes}, " f"layout={layout_id!r}/{checkpoint_spec.layout_id!r}" ) diff --git a/atom/model_ops/attentions/paged_state_copy.py b/atom/model_ops/attentions/paged_state_copy.py index 3ccab5497e..3907511c9b 100644 --- a/atom/model_ops/attentions/paged_state_copy.py +++ b/atom/model_ops/attentions/paged_state_copy.py @@ -5,8 +5,10 @@ from __future__ import annotations -from dataclasses import dataclass +from collections.abc import Sequence +from dataclasses import dataclass, field +import numpy as np import torch try: @@ -19,120 +21,306 @@ _TILE_BYTES = 4096 -@dataclass(frozen=True) -class ByteSegment: - ptr: int - num_bytes: int +@dataclass(frozen=True, eq=False) +class SegmentedCopyPlan: + """Where two ordered byte streams meet, in offsets rather than addresses. + Which source segment meets which destination segment, at what offset into + each and for how many bytes, follows from the two streams' *sizes* alone. + Addresses enter only when a copy is issued. -@dataclass(frozen=True) -class CopySpan: - src_ptr: int - dst_ptr: int - num_bytes: int + Holding them apart is what makes a finely segmented copy affordable. A + caller whose geometry outlives its copies — a checkpoint image is the same + shape for the life of the pool — walks the intersection once here and then + spends a few vector adds per copy where it used to spend a Python loop per + span. Measured on a DeepSeek-V4 image: 0.53 us a span against about none, + which is the difference between an image cut fine enough to save PAGE + units and one that costs more host time than the units are worth. + The first five arrays are parallel and one span long. `src` and `dst` name + the roles the plan was built in, not a direction: an intersection is + symmetric, so `write_descriptor` can read it either way and a restore + reuses the plan its store was cut by. -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()) + The last two are the same spans cut into tiles, which is the unit the + kernel actually runs on. Compared by identity (`eq=False`): the fields are + arrays, so a generated `__eq__` would raise rather than answer. + """ + + src_seg: np.ndarray + src_off: np.ndarray + dst_seg: np.ndarray + dst_off: np.ndarray + length: np.ndarray + # Which span a tile belongs to, and its byte offset inside that span. + span_of_tile: np.ndarray + tile_start: np.ndarray + # The two above, per device, uploaded on first use. Fixed geometry, so one + # upload serves every copy for the life of the plan. + _resident: dict = field(default_factory=dict, repr=False, compare=False) + + @property + def num_spans(self) -> int: + return int(self.length.size) + + @property + def num_tiles(self) -> int: + return int(self.span_of_tile.size) + + def tiling_on(self, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: + """`(span_of_tile, tile_start)` resident on `device`. + + Uploaded pageably, which synchronizes the stream — acceptable only + because it happens once per plan and the first copy is a warmup, not a + request. A caller that reaches this from a live batch pays the whole + outstanding queue for it. + """ + tables = self._resident.get(device) + if tables is None: + tables = ( + torch.from_numpy(self.span_of_tile).to(device), + torch.from_numpy(self.tile_start).to(device), + ) + self._resident[device] = tables + return tables + + def write_descriptor( + self, + out: np.ndarray, + src_bases: np.ndarray, + dst_bases: np.ndarray, + *, + forward: bool = True, + ) -> None: + """Fill `(copies * num_spans, 3)` int64 rows: source, destination, length. + + `src_bases` and `dst_bases` are `(copies, segments)` — one row of + segment addresses per copy, which is where a caller's geometry enters: + a slot's base plus a range's offset, a PAGE unit's base plus a + region's. `out` holds the copies back to back in the same order. + + Every copy in one call goes the same way, so a caller with both + directions makes two calls. `forward=False` copies the destination + stream back into the source instead. + + The copies are filled in one pass rather than one at a time. At these + sizes each of the three fills below is nearly all numpy call overhead + — a span table is a few hundred entries — so paying it per copy is + what made describing a batch a quarter of the whole copy path, once + the kernel stopped being the bottleneck. Batched it is about 7x + cheaper, and it is the same three lines with one more axis. + """ + # Numpy would broadcast a short `dst_bases` rather than complain, and + # every copy in the batch would then be aimed at the first image's + # addresses -- silent cross-request corruption, at raw pointers, with + # the other images' checkpoint records still claiming them. This is + # where the pointers are made, so it is where the shapes are checked; + # `launch_copy_descriptor` only inherits them. Both sides are asked + # the same way, down to the rank: checking only the leading axis of + # `dst_bases` would let a 1-D one reach `dst_bases[:, dst_seg]`, which + # raises an index error about an array the caller did not pass. + if ( + src_bases.ndim != 2 + or dst_bases.ndim != 2 + or len(dst_bases) != len(src_bases) + ): + raise ValueError( + f"a copy needs one row of bases per copy on both sides, got " + f"{src_bases.shape} and {dst_bases.shape}" + ) + copies = len(src_bases) + if out.shape != (copies * self.num_spans, 3): + raise ValueError( + f"a descriptor for {copies} copies of {self.num_spans} spans " + f"must be {(copies * self.num_spans, 3)}, got {out.shape}" + ) + src_col, dst_col = (0, 1) if forward else (1, 0) + rows = out.reshape(copies, self.num_spans, 3) + np.add(src_bases[:, self.src_seg], self.src_off, out=rows[:, :, src_col]) + np.add(dst_bases[:, self.dst_seg], self.dst_off, out=rows[:, :, dst_col]) + rows[:, :, 2] = self.length def plan_segmented_copy( - src: list[ByteSegment], - dst: list[ByteSegment], + src_sizes: Sequence[int], + dst_sizes: Sequence[int], total_bytes: int, -) -> list[CopySpan]: - """Intersect two ordered byte streams into physical copy spans.""" +) -> SegmentedCopyPlan: + """Intersect two ordered byte streams into the spans a copy is made of.""" 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: + if sum(src_sizes) < total_bytes: raise ValueError("source segmented stream is shorter than the copy") - if sum(s.num_bytes for s in dst) < total_bytes: + if sum(dst_sizes) < total_bytes: raise ValueError("destination segmented stream is shorter than the copy") - if any(s.num_bytes <= 0 for s in src + dst): + if any(size <= 0 for size in (*src_sizes, *dst_sizes)): raise ValueError("segmented streams cannot contain empty segments") - if total_bytes == 0: - return [] - spans: list[CopySpan] = [] + src_seg: list[int] = [] + src_off: list[int] = [] + dst_seg: list[int] = [] + dst_off: list[int] = [] + length: list[int] = [] src_i = dst_i = 0 - src_off = dst_off = 0 + src_used = dst_used = 0 remaining = total_bytes while remaining: - src_left = src[src_i].num_bytes - src_off - dst_left = dst[dst_i].num_bytes - dst_off + src_left = src_sizes[src_i] - src_used + dst_left = dst_sizes[dst_i] - dst_used nbytes = min(src_left, dst_left, remaining) - spans.append( - CopySpan( - src[src_i].ptr + src_off, - dst[dst_i].ptr + dst_off, - nbytes, - ) - ) + src_seg.append(src_i) + src_off.append(src_used) + dst_seg.append(dst_i) + dst_off.append(dst_used) + length.append(nbytes) remaining -= nbytes - src_off += nbytes - dst_off += nbytes - if src_off == src[src_i].num_bytes: + src_used += nbytes + dst_used += nbytes + if src_used == src_sizes[src_i]: src_i += 1 - src_off = 0 - if dst_off == dst[dst_i].num_bytes: + src_used = 0 + if dst_used == dst_sizes[dst_i]: dst_i += 1 - dst_off = 0 - return spans + dst_used = 0 + i64 = np.int64 + lengths = np.array(length, dtype=i64) + span_of_tile, tile_start = _tiling(lengths) + return SegmentedCopyPlan( + src_seg=np.array(src_seg, dtype=i64), + src_off=np.array(src_off, dtype=i64), + dst_seg=np.array(dst_seg, dtype=i64), + dst_off=np.array(dst_off, dtype=i64), + length=lengths, + span_of_tile=span_of_tile, + tile_start=tile_start, + ) + + +def _tiling(lengths: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + """Cut the spans into tiles: which span each is in, and where it starts. + + A copy kernel wants one program per tile *that exists*. Deriving that on + the device would need a search; deriving it here needs none, because a + plan's spans do not change. The result is two small arrays that go to the + device once and serve every copy the plan describes. + """ + counts = -(-lengths // _TILE_BYTES) + total = int(counts.sum()) + span_of_tile = np.repeat(np.arange(lengths.size, dtype=np.int32), counts) + # Exclusive prefix sum. Written as cumsum-minus-self rather than by + # prepending a zero and dropping the last, which has no answer for a plan + # with no spans at all: the prepended zero survives, and the subtraction + # below then fails to broadcast against an empty `counts`. + first = np.cumsum(counts) - counts + within = np.arange(total, dtype=np.int64) - np.repeat(first, counts) + return span_of_tile, within * _TILE_BYTES if triton is not None: @triton.jit def _copy_tiles_kernel( - src_ptrs, - dst_ptrs, - valid_bytes, - TILE_BYTES: tl.constexpr, + descriptor, + span_of_tile, + tile_start, + num_tiles, + num_spans, + TILE: 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) + # One program per tile that exists. The tiling is the plan's, resident + # on the device, so which span this tile belongs to and where it + # starts are two loads rather than a search. + # + # The two counts are ordinary arguments, not `tl.constexpr`: only + # `TILE` has to be one, for `tl.arange`. Specialising on the other two + # would key the compiled kernel to a pool geometry, so every image + # shape would miss the on-disk cache and pay the JIT again -- for a + # divide the copy does not notice. + # + # A row is `write_descriptor`'s: source, destination, length. Triton + # cannot read a plain module constant, so the three stay literal here + # and the round-trip test is what holds the two ends to the same order. + pid = tl.program_id(0) + op = pid // num_tiles + tile = pid % num_tiles + span = tl.load(span_of_tile + tile) + start = tl.load(tile_start + tile) + row = descriptor + (op * num_spans + span) * 3 + src_ptr = tl.load(row) + dst_ptr = tl.load(row + 1) + length = tl.load(row + 2) + offsets = start + tl.arange(0, TILE) + mask = offsets < length + src = (src_ptr.to(tl.int64) + offsets).to(tl.pointer_type(tl.uint8)) + dst = (dst_ptr.to(tl.int64) + offsets).to(tl.pointer_type(tl.uint8)) + tl.store(dst, tl.load(src, mask=mask), 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: +def launch_copy_descriptor(descriptor: torch.Tensor, plan: SegmentedCopyPlan) -> None: + """Copy every span a resident descriptor names, in one launch. + + One row per span, row-major, holding several copies back to back — every + op of a step goes out together. + + The upload is the caller's, and `descriptor` arrives on the device + already. That is not tidiness: a pageable `torch.from_numpy(x).to(dev)` + synchronizes the current stream, and this runs from `build()` with the + previous forward still enqueued, so the host waits out the whole queue + rather than the 800 KB. Measured behind 4 ms of work it cost 2.9 ms + against 0.1 ms staged through pinned memory — a cost the transfer's own + size says nothing about. `CpuGpuBuffer` is what the caller stages with, + and what the rest of this repo already uses to avoid exactly this. + + The grid is one program per tile that exists. It used to be rectangular, + `(spans, ceil(widest / TILE))`, which gives every span as many programs as + the *widest* one needs: on a DeepSeek-V4 image, whose spans run 8 KiB to + 1.4 MB, that is 46,364 programs to do 2,631 tiles of work. One op could + afford the waste; a batch cannot, and this path now always batches. At 256 + ops the rectangular grid measured 5.14 ms against 1.19 ms for this one. + + What makes the dense grid cheap is that the tiling is not a function of + the copy — the plan's spans are fixed, so `span_of_tile` and `tile_start` + are computed once and stay resident. No device-side search, and no + `widest` for a caller to get wrong: passing one too small used to truncate + every longer span silently, byte-correct on its prefix and stale on its + tail. + """ + rows = descriptor.shape[0] + if rows == 0: 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, + # Checked because this is where host arithmetic becomes device pointers: + # the kernel reads `descriptor` as a contiguous int64 (rows, 3) and indexes + # it by op, so a wrong shape or dtype is a wrong address rather than an + # error. + if descriptor.dtype != torch.int64 or descriptor.ndim != 2: + raise ValueError("a copy descriptor must be a 2-D int64 tensor") + if descriptor.shape[1] != 3 or not descriptor.is_contiguous(): + raise ValueError("a copy descriptor must be contiguous with 3 columns") + if rows % plan.num_spans: + raise ValueError( + f"descriptor of {rows} rows is not a whole number of " + f"{plan.num_spans}-span copies" + ) + span_of_tile, tile_start = plan.tiling_on(descriptor.device) + _copy_tiles_kernel[(rows // plan.num_spans * plan.num_tiles,)]( + descriptor, + span_of_tile, + tile_start, + plan.num_tiles, + plan.num_spans, + TILE=_TILE_BYTES, + # Four warps, not more: this kernel is nothing but load and store, and + # its speed turns out to be set by the width of one lane's access, + # `TILE / (num_warps * 64)`. Sixteen bytes is the fast point -- a + # 128-bit access -- and three unrelated (TILE, warps) pairs that land + # on it measured within 0.5% of each other, while eight warps halves + # the width and costs 12%. Raising this to fill more of the machine + # makes it slower, so it is not a knob to turn up. + num_warps=4, ) diff --git a/atom/model_ops/attentions/state_arena.py b/atom/model_ops/attentions/state_arena.py index c582f0fb5b..ecc7a2d369 100644 --- a/atom/model_ops/attentions/state_arena.py +++ b/atom/model_ops/attentions/state_arena.py @@ -35,7 +35,9 @@ at all: a field is one strided tensor, so it lands in one plane or the other. `plan_field_planes` decides which, `SplitStateArena` hides the split from consumers asking for a field by name, and what stays contiguous is a *slot* — -which is the range a checkpoint copies and a PD transfer registers anyway. +which is the range a PD transfer registers, and the range a checkpoint's own +is carved out of by `checkpoint_ranges_for`, since an image holds only the +fields a resumer reads (`StateField.in_checkpoint`). Backends stay in charge of what the fields are; this module only owns the arithmetic. The layout is deliberately the one DeepSeek-V4's PD staging path @@ -46,7 +48,9 @@ from __future__ import annotations import math +from collections.abc import Iterator from dataclasses import dataclass +from itertools import groupby import torch @@ -190,6 +194,18 @@ class StateField: # one of those wider rows, or the row index the kernel computes is off by a # fraction of a row and nothing about the view says so. align: int = 0 + # Whether a checkpoint image holds this field at all. A ring whose next + # reader starts exactly at the boundary the checkpoint was taken on owes + # one nothing: DeepSeek-V4's HCA compressor pools `[P, P + 128)` and a + # checkpoint sits on a multiple of 128, so every row its resumer reads is + # a row that same resumer writes. Declared by the field rather than worked + # out at copy time, so the copy path only reads what a field says about + # itself and the geometry stays in one place. + # + # All-or-nothing on purpose. A field carried in *some* of its rows is a + # ring, and which rows those are depends on the position the checkpoint + # was taken at — a phase this module is not given and must not guess. + in_checkpoint: bool = True def __post_init__(self): if self.align and self.align % _ALIGN: @@ -208,6 +224,22 @@ def bytes_per_entry(self) -> int: return self.layers * self.per_layer_numel * self.dtype.itemsize +def field_extents( + fields: list[StateField], +) -> Iterator[tuple[StateField, int, int]]: + """Each field with the `[start, end)` bytes it occupies in an entry. + + The one place the align-place-advance walk is written. An arena's field + offsets, the entry's own size and a checkpoint's ranges are three answers + to the same question and have to agree, so all three come from here. + """ + offset = 0 + for field in fields: + offset = _align_up(offset, max(_ALIGN, field.align)) + yield field, offset, offset + field.bytes_per_entry + offset += field.bytes_per_entry + + def entry_bytes_for(fields: list[StateField]) -> int: """Bytes one entry costs, including inter-field alignment. @@ -215,10 +247,35 @@ def entry_bytes_for(fields: list[StateField]) -> int: function rather than a property of a built arena — the byte budget and the allocation must come from the same expression or the two drift. """ - total = 0 - for field in fields: - total = _align_up(total, max(_ALIGN, field.align)) + field.bytes_per_entry - return _align_up(total) + end = 0 + for _, _, field_end in field_extents(fields): + end = field_end + return _align_up(end) + + +def checkpoint_ranges_for(fields: list[StateField]) -> list[tuple[int, int]]: + """`(offset, num_bytes)` of an entry a checkpoint image holds. + + Consecutive carried fields merge into one range, so the ordinary + all-carried case is a single range, and the alignment padding inside a run + rides along with it — splitting a range to shave padding costs more + descriptor than it saves. A field left out breaks the run, which is the + point: merging across it would put it back in the image. + """ + ranges: list[tuple[int, int]] = [] + for carried, run in groupby(field_extents(fields), lambda e: e[0].in_checkpoint): + if not carried: + continue + extents = list(run) + start = extents[0][1] + nbytes = extents[-1][2] - start + # A run of zero-byte fields spans nothing, and a zero-length range is + # not one: `plan_segmented_copy` refuses empty segments, and it is only + # reached on the first copy, so emitting one here would let a config + # size, cross-check and start cleanly and then abort mid-serving. + if nbytes: + ranges.append((start, nbytes)) + return ranges class StateArena: @@ -276,12 +333,7 @@ def __init__( if not 0 <= self.live_entries <= entries: raise ValueError(f"live_entries {self.live_entries} outside 0..{entries}") - offset = 0 - self._offsets: dict[str, int] = {} - for field in self.fields: - offset = _align_up(offset, max(_ALIGN, field.align)) - self._offsets[field.name] = offset - offset += field.bytes_per_entry + self._offsets = {f.name: start for f, start, _ in field_extents(self.fields)} self._by_name = {f.name: f for f in self.fields} # Zeroed, not `empty`: alignment padding falls outside every field diff --git a/atom/model_ops/attentions/v4_pool_geometry.py b/atom/model_ops/attentions/v4_pool_geometry.py index 6ddeea1d65..c8398254ec 100644 --- a/atom/model_ops/attentions/v4_pool_geometry.py +++ b/atom/model_ops/attentions/v4_pool_geometry.py @@ -64,6 +64,7 @@ from __future__ import annotations +from collections.abc import Iterable from dataclasses import dataclass # Compress ratios, as they appear in the model config's per-layer list. @@ -87,6 +88,40 @@ _KNOWN_RATIOS = frozenset(_ENTRY_ORDER) | {ABSENT_RATIO} +def merge_abutting(runs: Iterable[tuple[int, int]]) -> list[tuple[int, int]]: + """`(start, count)` pairs in order, with the ones that touch joined. + + Rows of a row space and bytes of a copy are both described this way, and + both want as few of them as possible: every range a checkpoint copy is cut + into costs a span, and a span costs grid. + + Ascending and disjoint is required, not assumed. Each run is compared only + against the one before it, so an out-of-order input merges nothing and + reads as legal: `[(0, 400), (256, 64), (320, 64)]` used to return + `[(0, 400), (256, 128)]`, which double-counts 128 bytes and overlaps the + first range. Nothing downstream can see that -- `checkpoint_image_bytes` + would over-count, and both the op validator and the sizing cross-check + compare against that same wrong number -- so the ordering is checked here + rather than left as a property of whoever happens to build the list. + """ + merged: list[tuple[int, int]] = [] + for start, count in runs: + if count < 0 or start < 0: + raise ValueError(f"a run must be non-negative, got ({start}, {count})") + if merged: + end = merged[-1][0] + merged[-1][1] + if start < end: + raise ValueError( + f"runs must ascend and not overlap: ({start}, {count}) " + f"starts inside the run ending at {end}" + ) + if start == end: + merged[-1] = (merged[-1][0], merged[-1][1] + count) + continue + merged.append((start, count)) + return merged + + def ring_offset_for(num_layers: int, ring_stride: int, ring_pos: int) -> int: """The layer-independent part of a window row, `f(q)`. @@ -210,6 +245,34 @@ def ring_row(self, layer_index: int, ring_pos: int) -> int: """Row of a window position, relative to the class's part of an entry.""" return layer_index * self.ring_stride + self.ring_offset(ring_pos) + def entry_row_runs(self) -> list[tuple[int, int]]: + """`(start, count)` runs of entry rows some `ring_row` reaches. + + The interleave is by ring position, not by layer: run `c` of the ring + owns the rows `[c*run_rows, (c+1)*run_rows)` and every layer of the + class has its positions for that run inside them. So all the whole + runs together are one contiguous range, and only the ring's last, + partial run is scattered — there each layer reaches + `ring_slots % ring_stride` rows of its own `ring_stride` slice and + leaves the rest. + + Those leftovers are `entry_rows` less `num_layers * ring_slots`: what + a layer-independent index formula costs (`entry_rows_for`). No + `(layer, position)` pair maps to one, so nothing writes or reads them, + which is what lets a copy that only has to preserve windows — a + checkpoint image is gathered back into a slot, never read by an + attention kernel — leave them out. + """ + run_rows = self.num_layers * self.ring_stride + whole, partial = divmod(self.ring_slots, self.ring_stride) + runs = [(0, whole * run_rows)] if whole else [] + if partial: + runs += [ + (whole * run_rows + i * self.ring_stride, partial) + for i in range(self.num_layers) + ] + return runs + class UnifiedPoolGeometry: """The row space, and every address formula that reads from it. @@ -515,6 +578,20 @@ def field_window_params( run_rows=self.ring_slots, ) + def entry_row_runs(self) -> list[tuple[int, int]]: + """`(start, count)` runs of entry rows any window reaches, in order. + + Every class's runs, offset into the entry and merged where they abut. + The complement is interleave padding — see `ClassLayout.entry_row_runs` + for why nothing reaches it. `classes` is built in entry order, so + walking it is walking the entry. + """ + return merge_abutting( + (cls.entry_offset + start, count) + for cls in self.classes.values() + for start, count in cls.entry_row_runs() + ) + def physical_slot(self, group: int) -> int: """Where pool group `group` sits in the plane. diff --git a/atom/plugin/vllm/deepseek_v4_bridge.py b/atom/plugin/vllm/deepseek_v4_bridge.py index 52ab8040fb..11d0e85419 100644 --- a/atom/plugin/vllm/deepseek_v4_bridge.py +++ b/atom/plugin/vllm/deepseek_v4_bridge.py @@ -136,11 +136,19 @@ def _v4_state_layout(vllm_config, kv_fp8: bool): torch.float32, float("-inf"), ), + # `in_checkpoint=False` for the same reason the native list says so + # (`deepseek_v4_attn._state_fields`): HCA pools `[P, P + 128)` with no + # overlap, so a checkpoint boundary owes it nothing. This list does not + # price an image today — only `plan_field_planes`, which ignores the + # flag, reads it — but two declarations of one layout disagreeing about + # the one rule `layout_id` fences is exactly what that fence cannot + # catch, since the id is derived from the native list alone. StateField( "hca_main_kv", n_hca, (128 + ring_extra, head_dim), torch.float32, + in_checkpoint=False, ), StateField( "hca_main_score", @@ -148,6 +156,7 @@ def _v4_state_layout(vllm_config, kv_fp8: bool): (128 + ring_extra, head_dim), torch.float32, float("-inf"), + in_checkpoint=False, ), ] row_widths = [head_dim * (1 if kv_fp8 else 2)] diff --git a/docs/scheduling_kv_cache_guide.md b/docs/scheduling_kv_cache_guide.md index 99af8cf5a3..d3ee476254 100644 --- a/docs/scheduling_kv_cache_guide.md +++ b/docs/scheduling_kv_cache_guide.md @@ -317,6 +317,14 @@ That number, `successor_room`, is mutability quantified. A rolling state (GDN re **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. +**A store takes what its image needs and nothing more.** `begin_store` asks for `units_per_checkpoint`, and `ensure_free_units` spends checkpoints only for the shortfall — free units come first, and `pop` hands out never-used blocks before cached ones, so a store reaches for the cache only once the pool has nothing spare. A store whose units are not reachable is dropped and counted in `checkpoints_dropped`, and it is dropped *before* evicting anything: eviction gives up only after it has emptied the cache, so an unreachable request would destroy the cache on its way to refusing. + +**There is no floor held back for live KV, and there cannot usefully be one.** `_fresh_block` raises when the pool is dry and nothing is evictable, so it is worth stating why the cache cannot take it there. A READY unpinned checkpoint is *already* available to live KV — `has_available_units` counts it and `ensure_free_units` will spend it — so the size of the cache is not the variable. What competes is the unevictable set, a checkpoint that is `COPYING` or held by a restore pin, and that set is confined to one pass: `schedule` publishes the previous batch's stores and releases its pins before it allocates anything, and this batch's stores are taken at batch construction, after every allocation. The one overlap is `allocate`, which pins a restore and then asks for fresh blocks in the same pass — and its own `can_allocate` counted that pin, protecting the checkpoint it is about to take and skipping the pins earlier admissions took. `may_append` never overlaps at all: the decode loop runs only in a pass that scheduled no prefill, so its pins were released at the top of it. Under contention the reachable outcome is therefore a refused admission, which the next pass retries, and never the raise. A floor would not improve on that in any case: live KV's demand is unbounded and legitimate — `allocate` takes a whole prompt's blocks, up to `max_model_len` of them — so no reserved quantity can promise it a block. *(The decode half of this rests on prefill and decode never sharing a pass; the mixed batch the scheduler has a TODO for would need the argument redone.)* + +**Eligibility and eviction policy are separate.** `_is_evictable` says whether a checkpoint may be spent — READY, unpinned, not the one the caller is protecting — and is the single rule `has_available_units` and `ensure_free_units` share. `_next_victim` says which eligible one to spend first, and is the only place the policy lives: least recently used today, one method to replace for another. `has_available_units` asks whether the eligible set reaches a count, which the order it is walked in cannot change -- only how soon the loop gets there -- so a new policy can change which checkpoint is spent and never what the gate answers. The refusal for an unreachable count sits in `ensure_free_units` itself rather than in a caller, because the bare loop gives up only after it has emptied the cache. + +The ladder asks a stricter question before it acts. A demand is an instruction to cut a prefill chunk onto a rung, and that cut costs the request a forward, so `_record_checkpoint_demand` records one only while `_checkpoint_has_room` holds — otherwise the forward is bought for a store `begin_store` is about to refuse. It is not "does an image fit": the admission asking takes its own block table first, so the question is `has_available_units(num_new_blocks + units_per_checkpoint)`, and it is asked with the same `protected_hash` `can_allocate` passes to `_has_page_units` on the next line, so the two gates of one pass agree on what eviction could reclaim. A pool with room for an image but not for this request *and* an image answers yes to the weaker question and refuses at `begin_store`, with the cut already bought. It is asked afresh on every attempt, because a demand recorded while the pool had room is not still affordable once it does not; what must not repeat is the counting, so the sequence carries a marker per counter rather than the gate reading a position it is about to overwrite. It remains a sample even so — the store happens many forwards later, at the rung this cut creates, against a pool that has moved — so what this gate removes is the loss knowable at admission and `checkpoints_dropped` counts the rest; the two are meant to be read together. What is *not* suppressed is the attribution: `num_wanted_hit_blocks` and hence `Lost-to-checkpoint` still say the reuse was declined for want of a checkpoint, because it was. `demands_declined_no_room` is where the difference shows, which keeps "the ladder is quiet because there is no demand" distinguishable from "the ladder is quiet because the pool is tight". + **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. Index order in the vacant half is not a fairness choice. Allocating lowest-first keeps the top of the pool cold — a high index is only reached at a concurrency high-water mark — which is what lets `retire_top` hand the pool's top group back when the KV/state boundary moves. When something *is* sitting there, `retire_top` relocates it and spends the least recently used checkpoint instead, wherever that one lives; retiring by index alone would be anti-LRU, since an index records the high-water mark at hand-out and is never refreshed by use. @@ -327,7 +335,9 @@ After pool sizing, the runner combines that capability with the optional `PagedS *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(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. +*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. + +An image holds part of a slot, not all of it (`PagedStateCheckpointSpec.image_bytes`, sized from `AttentionMetadataBuilder.checkpoint_image_bytes`). A resumer starts at the boundary the checkpoint was taken on, and a compressor that pools `ratio` tokens with no overlap begins its first pool exactly there — so every row it reads is one it writes, and the checkpoint owes it nothing. DeepSeek-V4's HCA state is 52% of a slot and entirely dead on that argument, which is what `StateField.in_checkpoint=False` declares; the sliding windows next to it are a sliding window, so every row a window position reaches is carried. The rows *between* them are not: a class interleaves its layers' windows so one index formula can serve every layer of it, and the rows that construction skips are reachable by no `(layer, position)` pair at all, so nothing writes or reads them (`UnifiedPoolGeometry.entry_row_runs`). Leaving them out costs 17.3% of the entry on the DSpark configuration and makes the image no longer a subsequence of the slot's rows — which is why a reader at the wrong version would gather every window row shifted. Both rules are named in `layout_id` (`nocopy=`, `entry=packed`) and fenced by its version, because two workers disagreeing about either would read one image at two layouts. It holds only while every ratio the model declares divides the quantity a checkpoint is aligned to -- `kv_cache_block_size * decode_context_parallel_size`, not `block_size` -- which `_assert_ratios_divide_the_alignment` enforces at startup against `hf_config.compress_ratios`. 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 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. diff --git a/tests/plugin/test_vllm_deepseek_v4_proxy_state_arena_layout.py b/tests/plugin/test_vllm_deepseek_v4_proxy_state_arena_layout.py index 9cac1606b0..9e66120eb1 100644 --- a/tests/plugin/test_vllm_deepseek_v4_proxy_state_arena_layout.py +++ b/tests/plugin/test_vllm_deepseek_v4_proxy_state_arena_layout.py @@ -48,11 +48,14 @@ def _flash_state_layout(kv_fp8: bool, *, ring_extra: int = 1): torch.float32, float("-inf"), ), + # Mirrors the bridge, `in_checkpoint` included: this list is that + # list's oracle, so a difference here is a difference nobody catches. StateField( "hca_main_kv", n_hca, (128 + ring_extra, head_dim), torch.float32, + in_checkpoint=False, ), StateField( "hca_main_score", @@ -60,6 +63,7 @@ def _flash_state_layout(kv_fp8: bool, *, ring_extra: int = 1): (128 + ring_extra, head_dim), torch.float32, float("-inf"), + in_checkpoint=False, ), ] row_widths = [head_dim * (1 if kv_fp8 else 2)] diff --git a/tests/test_page_unit_checkpoint.py b/tests/test_page_unit_checkpoint.py index cbbffa9a81..076a06972d 100644 --- a/tests/test_page_unit_checkpoint.py +++ b/tests/test_page_unit_checkpoint.py @@ -26,6 +26,7 @@ def make_store(num_units=20, unit_bytes=10, slot_bytes=25): page_unit_bytes=unit_bytes, slot_bytes=slot_bytes, layout_id="layout-v1", + image_bytes=slot_bytes, ), ) @@ -45,12 +46,13 @@ def ready(store, prefix_hash, src_slot=0): def test_runtime_spec_derives_units_and_has_a_minimal_wire_form(): - spec = PagedStateCheckpointSpec(10, 25, "layout-v1") + spec = PagedStateCheckpointSpec(10, 25, "layout-v1", image_bytes=25) assert spec.units_per_checkpoint == 3 assert spec.to_wire() == { "page_unit_bytes": 10, "slot_bytes": 25, + "image_bytes": 25, "layout_id": "layout-v1", } assert "units_per_checkpoint" not in spec.to_wire() @@ -62,12 +64,24 @@ def test_runtime_spec_derives_units_and_has_a_minimal_wire_form(): spec.slot_bytes = 30 +def test_units_are_priced_off_the_image_not_the_whole_slot(): + """An image holds part of a slot, so that part is what has to fit.""" + whole = PagedStateCheckpointSpec(10, 25, "layout-v1", image_bytes=25) + narrowed = PagedStateCheckpointSpec(10, 25, "layout-v1", image_bytes=11) + + assert whole.units_per_checkpoint == 3 + assert narrowed.units_per_checkpoint == 2 + + @pytest.mark.parametrize( "args", [ - (0, 25, "layout-v1"), - (10, -1, "layout-v1"), - (10, 25, ""), + (0, 25, "layout-v1", 25), + (10, -1, "layout-v1", 25), + (10, 25, "", 25), + (10, 25, "layout-v1", 0), + # An image cannot hold more than the slot it was taken from. + (10, 25, "layout-v1", 26), ], ) def test_runtime_spec_rejects_invalid_geometry(args): @@ -124,7 +138,7 @@ def test_empty_batch_does_not_complete_a_queued_restore(): pool = BlockPool(20) coordinator = PagedStateCheckpointCoordinator( pool, - PagedStateCheckpointSpec(10, 25, "layout-v1"), + PagedStateCheckpointSpec(10, 25, "layout-v1", image_bytes=25), enabled=True, ) checkpoint_id, _ = ready(coordinator.store, 101) @@ -204,3 +218,165 @@ def test_clear_releases_ready_images_but_defers_a_pinned_reader(): store.complete_inflight() assert not store.records assert pool.num_free == 20 + + +def _filled(num_units, unit_bytes, image_bytes, count): + """A store holding `count` READY checkpoints, oldest first.""" + pool = BlockPool(num_units) + store = PageUnitCheckpointStore( + pool, + PagedStateCheckpointSpec( + page_unit_bytes=unit_bytes, + slot_bytes=image_bytes, + layout_id="layout-v1", + image_bytes=image_bytes, + ), + ) + for prefix_hash in range(count): + assert store.begin_store(prefix_hash, src_slot=0) is not None + store.complete_inflight() + return pool, store + + +def test_a_store_with_free_units_spends_no_checkpoint(): + """Free units first. A store asking for what is already there evicts nothing. + + The cache is not a reservoir a store drains to a level -- it takes its own + image's worth. This used to be `needed + reserve_units`, which meant an + accepted store spent tens of checkpoints to build a cushion for live KV + that live KV never needed. + """ + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=100, count=3) + assert pool.num_free == 70 + + assert store.begin_store(999, src_slot=0) is not None + + assert store.evictions == 0, "a store with 70 free units spent a checkpoint" + assert len(store.records) == 4, "the cache lost an entry it did not have to" + + +def test_a_store_spends_only_the_shortfall(): + """Short by half an image: one checkpoint covers it, and only one goes.""" + pool, store = _filled(num_units=35, unit_bytes=10, image_bytes=100, count=3) + assert pool.num_free == 5, "the pool is meant to be short by half an image" + + assert store.begin_store(999, src_slot=0) is not None + + assert store.evictions == 1, "the shortfall cost more than one checkpoint" + assert store.lookup(0) < 0, "the victim was not the oldest" + assert store.lookup(1) >= 0 and store.lookup(2) >= 0 + + +def test_a_dropped_store_evicts_nothing(): + """A store that cannot get its units has to cost nothing. + + `ensure_free_units` gives up only after it has evicted everything it can, + so asking it for units that are not there would destroy the cache on the + way to refusing. `begin_store` asks whether they are reachable first. + """ + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=10, count=50) + # Live KV takes every unit the checkpoints left. + pool.reserve_units(pool.num_free, ("live-kv", 0)) + for record in store.records.values(): + record.pin_count = 1 # every checkpoint is being read, so none is spendable + + assert store.begin_store(999, src_slot=0) is None + + assert store.evictions == 0, "a dropped store evicted" + assert len(store.records) == 50, "a dropped store cost the cache" + + +def test_the_eviction_policy_cannot_move_the_gate(): + """Eligibility is shared; order is policy. Only the second one may change. + + `has_available_units` asks whether the eligible set reaches a count, which + the order it is walked in cannot change -- only how soon the loop gets + there. Swapping the policy here has to leave every gate answer identical. + """ + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=10, count=6) + lru_pick = store._next_victim() + available = [store.has_available_units(n) for n in range(0, 101, 10)] + + def newest_first(protected=-1): + return next( + ( + cid + for cid in reversed(store._lru) + if store._is_evictable(cid, protected) + ), + -1, + ) + + store._next_victim = newest_first + + assert [store.has_available_units(n) for n in range(0, 101, 10)] == available + assert store._next_victim() != lru_pick, "the policy swap did not take" + + pool.reserve_units(pool.num_free, ("live-kv", 0)) + assert store.ensure_free_units(1) + assert store.lookup(5) < 0, "the new policy's victim was not spent" + assert store.lookup(0) >= 0, "the LRU victim was spent under another policy" + + +def test_a_store_still_recycles_the_oldest_checkpoint(): + """The gate refuses a store; it does not stop the policy doing its job.""" + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=10, count=100) + assert pool.num_free == 0 + + assert store.begin_store(999, src_slot=0) is not None + + assert store.evictions == 1 + assert store.lookup(0) < 0, "the victim was not the oldest" + + +def test_a_restore_takes_no_units(): + """Only new images need units; reading one back does not.""" + pool = BlockPool(20) + spec = PagedStateCheckpointSpec(10, 25, "layout-v1", image_bytes=25) + store = PageUnitCheckpointStore(pool, spec) + ready(store, 101) + pool.reserve_units(pool.num_free, ("live-kv", 0)) + + assert store.begin_restore(101, dst_slot=4) is not None + + +def test_an_unreachable_count_evicts_nothing_whoever_asks(): + """The refusal lives in `ensure_free_units`, not in one of its callers. + + `begin_store` used to carry the reachability test itself, which left + `BlockManager._ensure_page_units` calling the raw loop -- harmless only + because its single caller passes 1, where there is nothing to spend before + giving up. Ask for more than the cache can reach and the bare loop empties + it and refuses anyway, which is the behaviour 0c46f4ed3 removed from one + call site and left available at the other. + """ + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=10, count=50) + pool.reserve_units(pool.num_free, ("live-kv", 0)) + assert pool.num_free == 0 and len(store.records) == 50 + + # 50 spendable units against a request for 60: unreachable, and reachable + # only after spending every one of them. + assert not store.ensure_free_units(60) + + assert store.evictions == 0, "a refused request emptied the cache" + assert len(store.records) == 50 + + +def test_a_store_refuses_when_the_policy_leaves_the_loop_short(): + """`begin_store` reads the answer rather than assuming it. + + `_next_victim` exists to be replaced. A policy that passes over an + eligible checkpoint makes the loop end short of `count`, and a + `begin_store` that assumed success would take an identity for a store that + cannot happen -- the record is safe only because `pool.reserve_units` + happens to refuse second. + """ + pool, store = _filled(num_units=100, unit_bytes=10, image_bytes=10, count=100) + assert pool.num_free == 0 + store._next_victim = lambda protected=-1: -1 # a policy that spends nothing + before = store._next_checkpoint_id + + assert store.begin_store(999, src_slot=0) is None + + assert store._next_checkpoint_id == before, "a refused store took an identity" + assert len(store.records) == 100 diff --git a/tests/test_paged_state_copy_planner.py b/tests/test_paged_state_copy_planner.py index 077849447e..155baa75cd 100644 --- a/tests/test_paged_state_copy_planner.py +++ b/tests/test_paged_state_copy_planner.py @@ -1,22 +1,39 @@ # SPDX-License-Identifier: MIT +"""Cutting one segmented byte stream against another, and issuing the result. + +The plan is addressless on purpose (`SegmentedCopyPlan`), so these tests are in +two halves: that the cut lands where it should, and that feeding it a pair of +base-address vectors reconstitutes the copy those cuts describe — including +backwards, which is what a restore is. +""" + +import numpy as np import pytest import torch from atom.model_ops.attentions.paged_state_copy import ( - ByteSegment, - launch_copy_spans, + launch_copy_descriptor, 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)] +def describe(plan, src_bases, dst_bases, forward=True): + """The plan at concrete addresses, as `(src, dst, length)` triples.""" + out = np.empty((plan.num_spans, 3), dtype=np.int64) + plan.write_descriptor( + out, + np.array([src_bases], dtype=np.int64), + np.array([dst_bases], dtype=np.int64), + forward=forward, + ) + return [tuple(int(x) for x in row) for row in out] - spans = plan_segmented_copy(src, dst, total_bytes=12) - assert [(s.src_ptr, s.dst_ptr, s.num_bytes) for s in spans] == [ +def test_segmented_stream_intersection_preserves_wire_order(): + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + + assert describe(plan, [1000, 2000], [3000, 4000, 5000]) == [ (1000, 3000, 3), (1003, 4000, 2), (2000, 4002, 2), @@ -24,13 +41,91 @@ def test_segmented_stream_intersection_preserves_wire_order(): ] +def test_a_reversed_descriptor_is_the_same_cut_the_other_way(): + """A restore reuses its store's plan, so the two must mirror exactly.""" + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + + forward = describe(plan, [1000, 2000], [3000, 4000, 5000]) + backward = describe(plan, [1000, 2000], [3000, 4000, 5000], forward=False) + + assert backward == [(dst, src, n) for src, dst, n in forward] + + +def test_the_plan_does_not_depend_on_the_addresses(): + """The same cut at two sets of bases differs by exactly the bases. + + Asserting only that the *unchanged* columns stayed put would pass for a + `write_descriptor` that ignored `src_bases` altogether, so the moved + column is checked against the delta rather than merely for inequality. + """ + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + delta = 999_000 + + here = describe(plan, [1000, 2000], [3000, 4000, 5000]) + moved = describe(plan, [1000 + delta, 2000], [3000, 4000, 5000]) + + assert [n for _, _, n in here] == [n for _, _, n in moved] + assert [d for _, d, _ in here] == [d for _, d, _ in moved] + # Only spans out of source segment 0 move, and each by exactly `delta`. + assert [s for s, _, _ in moved] == [ + src + (delta if seg == 0 else 0) + for (src, _, _), seg in zip(here, plan.src_seg, strict=True) + ] + + 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 + plan = plan_segmented_copy([13], [5, 5, 5], total_bytes=13) + spans = describe(plan, [1000], [2000, 3000, 4000]) + + assert sum(n for _, _, n in spans) == 13 + assert spans[-1][1] == 4000 + assert spans[-1][2] == 3 + + +def test_the_tiling_covers_every_span_exactly_once(): + """The grid is one program per tile, so the tiling is the copy's extent.""" + plan = plan_segmented_copy([13], [5, 5, 5], total_bytes=13) + + assert plan.num_spans == 3 + # Each of the three spans is under one tile, so one tile each, all at + # offset zero inside their span. + assert plan.num_tiles == 3 + assert list(plan.span_of_tile) == [0, 1, 2] + assert list(plan.tile_start) == [0, 0, 0] + + +def test_a_long_span_is_cut_into_consecutive_tiles(): + plan = plan_segmented_copy([10_000], [10_000], total_bytes=10_000) + + assert plan.num_spans == 1 + assert plan.num_tiles == 3 # 4096 + 4096 + 1808 + assert list(plan.span_of_tile) == [0, 0, 0] + assert list(plan.tile_start) == [0, 4096, 8192] + + +def test_the_tiling_reaches_the_end_of_every_span(): + """Nothing is dropped off a span's tail, whatever the sizes are.""" + plan = plan_segmented_copy([9_000, 300, 20_000], [29_300], total_bytes=29_300) + + covered = {} + for span, start in zip(plan.span_of_tile, plan.tile_start, strict=True): + covered.setdefault(int(span), []).append(int(start)) + for span, starts in covered.items(): + assert starts == list(range(0, int(plan.length[span]), 4096)) + + +@pytest.mark.parametrize( + "src, dst, total, message", + [ + ([5], [5], -1, "non-negative"), + ([3], [5], 5, "source segmented stream is shorter"), + ([5], [3], 5, "destination segmented stream is shorter"), + ([5, 0], [5], 5, "cannot contain empty segments"), + ], +) +def test_an_impossible_copy_is_refused(src, dst, total, message): + with pytest.raises(ValueError, match=message): + plan_segmented_copy(src, dst, total) @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a GPU") @@ -40,22 +135,145 @@ def test_descriptor_kernel_round_trips_random_bytes_with_partial_tail(): 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(), + units = [4096, 4096, image.numel() - 8192] + unit_bases = np.array( + [image.data_ptr(), image.data_ptr() + 4096, image.data_ptr() + 8192], + dtype=np.int64, ) - launch_copy_spans(gather, device) + plan = plan_segmented_copy([original.numel()], units, original.numel()) + + for slot_ptr, forward in ( + (original.data_ptr(), True), + (restored.data_ptr(), False), + ): + descriptor = np.empty((plan.num_spans, 3), dtype=np.int64) + plan.write_descriptor( + descriptor, + np.array([[slot_ptr]], dtype=np.int64), + unit_bases[None], + forward=forward, + ) + launch_copy_descriptor(torch.from_numpy(descriptor).to(device), plan) 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) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a GPU") +def test_several_copies_ride_in_one_descriptor(): + """Production batches every op of a step into a single launch.""" + device = torch.device("cuda") + sources = [ + torch.randint(0, 256, (5_000,), dtype=torch.uint8, device=device) + for _ in range(3) + ] + images = [torch.zeros(6_000, dtype=torch.uint8, device=device) for _ in range(3)] + plan = plan_segmented_copy([5_000], [2_048, 2_048, 1_904], 5_000) + + descriptor = np.empty((3 * plan.num_spans, 3), dtype=np.int64) + for i, (src, image) in enumerate(zip(sources, images, strict=True)): + plan.write_descriptor( + descriptor[i * plan.num_spans : (i + 1) * plan.num_spans], + np.array([[src.data_ptr()]], dtype=np.int64), + np.array( + [ + [ + image.data_ptr(), + image.data_ptr() + 2_048, + image.data_ptr() + 4_096, + ] + ], + dtype=np.int64, + ), + ) + launch_copy_descriptor(torch.from_numpy(descriptor).to(device), plan) + + torch.cuda.synchronize() + for src, image in zip(sources, images, strict=True): + assert torch.equal(image[:5_000], src) + + +def test_a_batch_describes_each_copy_the_way_one_call_would(): + """The vectorised fill has to be indistinguishable from a loop. + + `write_descriptor` fills every copy in one pass because doing it per copy + was a quarter of the whole copy path. That is only allowed to be faster, + so the batch is checked against the copies written one at a time -- an + axis dropped or an offset broadcast the wrong way would otherwise be a + descriptor full of plausible addresses. + """ + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + src_bases = np.array([[1000, 2000], [1100, 2100], [1200, 2200]], dtype=np.int64) + dst_bases = np.array( + [[7000, 8000, 9000], [7100, 8100, 9100], [7200, 8200, 9200]], dtype=np.int64 + ) + + for forward in (True, False): + batched = np.empty((3 * plan.num_spans, 3), dtype=np.int64) + plan.write_descriptor(batched, src_bases, dst_bases, forward=forward) + + one_at_a_time = np.empty_like(batched) + for i in range(3): + plan.write_descriptor( + one_at_a_time[i * plan.num_spans : (i + 1) * plan.num_spans], + src_bases[i : i + 1], + dst_bases[i : i + 1], + forward=forward, + ) + + assert np.array_equal(batched, one_at_a_time), f"forward={forward}" + + +def test_a_zero_byte_copy_is_an_empty_plan_not_a_broadcast_error(): + """0 is a legal `total_bytes` by this function's own contract. + + It validates negatives and empty segments, so a caller reads 0 as allowed + -- and the tiling used to build its exclusive prefix sum by prepending a + zero, which has no answer when there are no spans and failed to broadcast + instead. An empty plan is the answer; `launch_copy_descriptor` returns on + a zero-row descriptor before it can divide by `num_spans`. + """ + plan = plan_segmented_copy([5], [5], total_bytes=0) + + assert plan.num_spans == 0 + assert plan.num_tiles == 0 + + +def test_bases_that_do_not_cover_every_copy_are_refused(): + """Numpy would broadcast them, and every copy would go to image zero. + + Verified before the guard: a one-row `dst_bases` against three copies + produced three identical destination rows -- three GPU copies into the + same image, no exception, two images left holding stale bytes their + checkpoint records still claim. This is where host arithmetic becomes a + raw pointer, so it is where the shape is checked. + """ + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + out = np.empty((3 * plan.num_spans, 3), dtype=np.int64) + three = np.zeros((3, 2), dtype=np.int64) + + with pytest.raises(ValueError, match="one row of bases per copy"): + plan.write_descriptor(out, three, np.zeros((1, 3), dtype=np.int64)) + + with pytest.raises(ValueError, match="must be"): + plan.write_descriptor( + np.empty((2, 3), dtype=np.int64), three, np.zeros((3, 3), dtype=np.int64) + ) + + +def test_a_flat_dst_bases_is_refused_by_this_guard_not_by_numpy(): + """Both sides are asked the same way, down to the rank. + + A guard that checked only the leading axis of `dst_bases` accepted a 1-D + one -- three entries, three copies, the shapes agree -- and it then failed + two lines later inside `dst_bases[:, dst_seg]`, raising an index error + about an array the caller never passed. Same rejection, unreadable reason. + """ + plan = plan_segmented_copy([5, 7], [3, 4, 5], total_bytes=12) + out = np.empty((3 * plan.num_spans, 3), dtype=np.int64) + three = np.zeros((3, 2), dtype=np.int64) + + with pytest.raises(ValueError, match="one row of bases per copy"): + plan.write_descriptor(out, three, np.zeros(3, dtype=np.int64)) diff --git a/tests/test_state_arena.py b/tests/test_state_arena.py index 23bc8c8e42..357f643457 100644 --- a/tests/test_state_arena.py +++ b/tests/test_state_arena.py @@ -10,6 +10,7 @@ """ import math +from dataclasses import replace from itertools import pairwise import pytest @@ -18,7 +19,9 @@ from atom.model_ops.attentions.state_arena import ( StateArena, StateField, + checkpoint_ranges_for, entry_bytes_for, + field_extents, plan_field_planes, plan_regions, ) @@ -40,6 +43,16 @@ def build(fields=V4_LIKE, entries=5) -> StateArena: return StateArena(fields, entries, device="cpu") +def carried_bytes(fields) -> int: + """Bytes of an entry a checkpoint image holds, the long way round. + + Spelled out here rather than imported: the module used to export this and + nothing but these tests called it, and a helper kept alive by its own + tests is not an interface. + """ + return sum(nbytes for _, nbytes in checkpoint_ranges_for(fields)) + + class TestEntryBytes: def test_sum_of_fields_when_naturally_aligned(self): @@ -413,3 +426,107 @@ def test_field_order_inside_a_plane_is_the_declared_one(self): def test_a_row_space_with_no_planes_is_rejected(self): with pytest.raises(ValueError, match="at least one plane"): plan_field_planes(V4_LIKE, []) + + +class TestCheckpointRanges: + """Which bytes of an entry a checkpoint image holds. + + A field declaring `in_checkpoint=False` is dead at a checkpoint boundary — + the resumer writes every row of it before reading any — so the image must + not carry it, and must not carry its neighbours' padding by accident + either. The flag is a bool because "some of its rows" would depend on the + position the checkpoint was taken at, which this module is never given. + """ + + @staticmethod + def without_hca(): + return [ + replace(f, in_checkpoint=False) if f.name.startswith("hca_") else f + for f in V4_LIKE + ] + + def test_an_all_carried_entry_is_one_range(self): + assert checkpoint_ranges_for(V4_LIKE) == [(0, entry_bytes_for(V4_LIKE))] + assert carried_bytes(V4_LIKE) == entry_bytes_for(V4_LIKE) + + def test_a_dropped_field_is_not_in_the_image(self): + fields = self.without_hca() + arena = StateArena(fields, 5, device="cpu") + + (start, nbytes), *rest = checkpoint_ranges_for(fields) + assert not rest, "the four CSA fields are adjacent, so they merge" + assert start == 0 + # Stops at the first dropped field rather than running to entry_bytes. + assert nbytes <= arena.field_offset("hca_main_kv") + assert carried_bytes(fields) < entry_bytes_for(fields) + + def test_a_dropped_field_breaks_the_run_it_sits_in(self): + """Merging across it would put it back in the image.""" + fields = [ + StateField("a", 1, (4, 8), torch.float32), + StateField("dead", 1, (64, 8), torch.float32, in_checkpoint=False), + StateField("b", 1, (4, 8), torch.float32), + ] + arena = StateArena(fields, 2, device="cpu") + dead_start = arena.field_offset("dead") + dead_end = dead_start + fields[1].bytes_per_entry + + ranges = checkpoint_ranges_for(fields) + + assert [start for start, _ in ranges] == [ + arena.field_offset("a"), + arena.field_offset("b"), + ] + for start, nbytes in ranges: + assert start >= dead_end or start + nbytes <= dead_start + + def test_the_image_spans_the_carried_run_padding_included(self): + """From the first carried field's start to the last one's end. + + Derived from `field_extents` rather than from the function under + test, and not from `sum(bytes_per_entry)` either: the alignment + between two carried fields rides along, because splitting a range to + shave it costs more descriptor than it saves. + """ + fields = self.without_hca() + carried = [(s, e) for f, s, e in field_extents(fields) if f.in_checkpoint] + + assert carried_bytes(fields) == carried[-1][1] - carried[0][0] + assert carried_bytes(fields) >= sum( + f.bytes_per_entry for f in fields if f.in_checkpoint + ) + + def test_the_ranges_land_where_the_arena_put_the_fields(self): + """Sizing and the copy have to read one layout, not two.""" + fields = self.without_hca() + arena = StateArena(fields, 5, device="cpu") + + carried = [f for f in fields if f.in_checkpoint] + (start, nbytes), *rest = checkpoint_ranges_for(fields) + assert not rest + assert start == arena.field_offset(carried[0].name) + last = carried[-1] + assert start + nbytes == arena.field_offset(last.name) + last.bytes_per_entry + + def test_a_wholly_dropped_entry_has_no_ranges(self): + fields = [StateField("dead", 1, (8, 8), torch.float32, in_checkpoint=False)] + + assert checkpoint_ranges_for(fields) == [] + assert carried_bytes(fields) == 0 + + def test_a_zero_byte_carried_run_is_not_a_range(self): + """A range of no bytes is refused downstream, and only on first use. + + `plan_segmented_copy` rejects empty segments, and it is reached lazily + on the first checkpoint copy -- so a field list that produced one would + size, cross-check and start cleanly, then abort mid-serving on the + first request to cross a rung. + """ + empty = StateField("no_layers", 0, (4, 4), torch.float32) + dead = StateField("dead", 1, (8, 8), torch.float32, in_checkpoint=False) + + assert checkpoint_ranges_for([empty]) == [] + # Between two dropped fields it is a run of its own, so nothing merges + # it away either. + assert checkpoint_ranges_for([dead, empty, dead]) == [] + assert all(n > 0 for _, n in checkpoint_ranges_for([empty, *V4_LIKE])) diff --git a/tests/test_state_checkpoint.py b/tests/test_state_checkpoint.py index 2cf28337d4..e7bd87ed4f 100644 --- a/tests/test_state_checkpoint.py +++ b/tests/test_state_checkpoint.py @@ -35,7 +35,7 @@ BLOCK = 4 MIN_FORK = 8 -PAGED_COPY_SPEC = PagedStateCheckpointSpec(10, 25, "test-layout-v1") +PAGED_COPY_SPEC = PagedStateCheckpointSpec(10, 25, "test-layout-v1", image_bytes=25) 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) @@ -1364,7 +1364,7 @@ def test_a_copy_never_asks_the_resumer_for_room(self): forking = StateGroupPool(4, StateTransfer.fork(4), hash_block_size=1) copying = PagedStateCheckpointCoordinator( BlockPool(4), - PagedStateCheckpointSpec(1, 1, "test-layout"), + PagedStateCheckpointSpec(1, 1, "test-layout", image_bytes=1), enabled=True, ) assert isinstance(copying, StateCache) @@ -1438,6 +1438,17 @@ def test_order_between_classes_does_not_change_the_answer(self): INTERVAL = 4 * BLOCK PROMPT = list(range(44)) # 11 blocks; last never reused, so 10 are hittable +# An image that costs more units than a request's blocks do. That is the shape +# where the ladder's question and admission's question can disagree: the pool +# still has room for the request and not for the checkpoint. With an image the +# size of a couple of blocks the two run out together and there is nothing to +# test. +BIG_IMAGE_SPEC = PagedStateCheckpointSpec(10, 400, "test-layout-big", image_bytes=400) +BIG_IMAGE_RUNTIME = StateRuntime( + transfer=StateTransfer.copy(BIG_IMAGE_SPEC.layout_id), + checkpoint_spec=BIG_IMAGE_SPEC, +) + def demand_config(**overrides): """A grid too coarse to cover the prompt, so demand has room to show. @@ -1452,6 +1463,17 @@ def demand_config(**overrides): return ckpt_config(**overrides) +def an_image_fits_on_its_own(checkpoints) -> bool: + """What the demand gate used to ask, kept as the contrast it is read against. + + The gate now asks whether an image fits *after* the admission has taken + its own blocks, and the tests below turn on the pool state where the two + answers differ. Written out here rather than left as a method on the + store, which would be a production API nothing in production asks. + """ + return checkpoints.has_available_units(checkpoints.store.units_per_checkpoint) + + class TestDemandDrivenCheckpoints: """A rung placed where a request was seen to want one. @@ -1496,6 +1518,185 @@ def test_the_third_request_finds_what_the_second_was_missing(self): assert third.checkpoint_demand_pos == 0 # nothing left to want assert forward_on_the_ladder(bm, third) == [] + def test_a_demand_the_floor_would_refuse_is_not_recorded(self): + """A cut costs a forward; buying one for a refused store is pure loss. + + `begin_store` drops a checkpoint whose units are not reachable. + Recording a demand for it anyway would still shorten the + request's prefill chunk, so the ladder asks the same question the + store will — and the reuse attribution is unaffected, because that + reuse really was declined for want of a checkpoint. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) + checkpoints = bm.paged_state_checkpoints + assert an_image_fits_on_its_own(checkpoints) + + # Live KV takes the pool down to where an image no longer fits but the + # resumer's own blocks still do -- the state under real pressure, where + # admission goes through and only the store cannot. + spare = -(-len(PROMPT) // BLOCK) + 1 + assert spare < checkpoints.store.units_per_checkpoint + bm.kv.reserve_units(bm.kv.num_free - spare, ("live-kv", 0)) + assert not an_image_fits_on_its_own(checkpoints) + + second = stateful_seq(PROMPT) + hit = bm.can_allocate(second) + bm.allocate(second, hit) + + # The reuse is still attributed to a missing checkpoint... + assert second.num_wanted_hit_blocks > hit, "the attribution was suppressed too" + # ...but nothing is cut for a store that would be refused. + assert second.checkpoint_demand_pos == 0, "a refused store still cut a chunk" + funnel = bm.checkpoint_funnel() + assert funnel["demands_declined_no_room"] == 1 + assert funnel["demands_recorded"] == 0 + assert funnel["chunks_cut_for_demand"] == 0 + assert forward_on_the_ladder(bm, second) == [32], "the grid rung, no demand cut" + + def _tighten_past_an_image(self, bm): + """Leave room for a resumer's blocks but not for a checkpoint image.""" + checkpoints = bm.paged_state_checkpoints + spare = -(-len(PROMPT) // BLOCK) + 1 + assert spare < checkpoints.store.units_per_checkpoint + bm.kv.reserve_units(bm.kv.num_free - spare, ("live-kv", 0)) + assert not an_image_fits_on_its_own(checkpoints) + + def test_a_demand_is_refused_when_the_admission_itself_drains_the_pool(self): + """The blocks this request takes come first, so they count. + + A pool with room for an image but not for the request *and* the image + answers yes to "does an image fit" -- and then the admission takes its + block table, `begin_store` refuses many forwards later, and the cut + this gate exists to withhold has already been bought. The funnel shows + nothing, because the decline happened somewhere that does not count. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) + checkpoints = bm.paged_state_checkpoints + image = checkpoints.store.units_per_checkpoint + blocks = -(-len(PROMPT) // BLOCK) + + # Enough for an image on its own, not for this request and an image. + bm.kv.reserve_units(bm.kv.num_free - (image + blocks // 2), ("live-kv", 0)) + assert an_image_fits_on_its_own(checkpoints), "the old question still says yes" + + second = stateful_seq(PROMPT) + assert bm.can_allocate(second) >= 0, "admission itself must still go through" + + assert second.checkpoint_demand_pos == 0, "bought a cut the store cannot use" + assert bm.checkpoint_funnel()["demands_declined_no_room"] == 1 + + def test_both_gates_in_one_pass_protect_the_same_checkpoint(self): + """`_checkpoint_has_room` and `_has_page_units` agree on what is spendable. + + The second excludes the checkpoint this admission is about to pin -- + it is about to be read, so eviction cannot have it. The first used to + count it as reclaimable, so with the pool resting on exactly that one + image the two gates in a single pass gave opposite answers. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + first = stateful_seq(PROMPT) + published = publish_at_boundary(bm, first) + bm.take_state_maintenance_ops() + bm.complete_previous_state_batch() + checkpoints = bm.paged_state_checkpoints + assert checkpoints.store.contains(published) + + # Nothing spare: the only spendable units are that one checkpoint's. + bm.kv.reserve_units(bm.kv.num_free, ("live-kv", 0)) + + assert bm._checkpoint_has_room(0, protected_hash=None), "the setup is wrong" + assert not bm._checkpoint_has_room( + 0, protected_hash=published + ), "the checkpoint about to be pinned was counted as spendable" + + def test_the_checkpoint_being_resumed_from_is_not_counted_as_spendable(self): + """Through `can_allocate`, where the two gates actually meet. + + The seq hits a checkpoint and wants a further one, so the pin and the + demand happen in the same call. Rest the pool on exactly that one + image and the answer turns on whether the gate knows it is spoken for: + counting it leaves the ladder cutting a chunk for a store that has no + units left to take. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + first = stateful_seq(PROMPT) + run_prompt_on_the_ladder(bm, first) + bm.take_state_maintenance_ops() + bm.complete_previous_state_batch() + image = bm.paged_state_checkpoints.store.units_per_checkpoint + assert len(bm.paged_state_checkpoints.store.records) == 1, "one image only" + + second = stateful_seq(PROMPT) + # Leave the request's own blocks plus half an image: reachable only by + # spending the very checkpoint `second` is about to resume from. + spare = -(-len(PROMPT) // BLOCK) + image // 2 + bm.kv.reserve_units(bm.kv.num_free - spare, ("live-kv", 0)) + + hit = bm.can_allocate(second) + + assert hit > 0, "the seq is supposed to resume from that checkpoint" + assert second.num_wanted_hit_blocks > hit, "and to want a further one" + assert second.checkpoint_demand_pos == 0, "spent an image already spoken for" + assert bm.checkpoint_funnel()["demands_declined_no_room"] == 1 + + def test_a_demand_recorded_while_there_was_room_is_withdrawn_when_it_goes(self): + """The gate is the store's question, so it has to be asked afresh. + + `can_allocate` re-runs for a sequence the queue keeps deferring. One + that recorded a demand while the pool had room, and is then re-admitted + against a pool that does not, is exactly the case the gate exists for: + the cut it would buy is now pure loss. Reading the position the gate is + about to overwrite made it a one-shot and let that cut through. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) + second = stateful_seq(PROMPT) + assert bm.can_allocate(second) >= 0 + recorded = second.checkpoint_demand_pos + assert recorded, "the first attempt was supposed to record a demand" + + self._tighten_past_an_image(bm) + bm.can_allocate(second) + + assert second.checkpoint_demand_pos == 0, "a stale answer bought the cut" + assert bm.checkpoint_funnel()["demands_declined_no_room"] == 1 + + def test_a_deferred_sequence_is_counted_once_however_often_it_asks(self): + """One request under pressure, not one per admission attempt. + + `demands_declined_no_room` is read against `demands_recorded`, so a + counter that fires per attempt makes the funnel unreadable under the + only pressure anyone reads it in -- and a decline writes 0 into the + position, so the position cannot be the marker that stops it. + """ + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) + self._tighten_past_an_image(bm) + + second = stateful_seq(PROMPT) + for _ in range(5): + bm.can_allocate(second) + + funnel = bm.checkpoint_funnel() + assert funnel["demands_declined_no_room"] == 1, "counted per attempt" + assert funnel["demands_recorded"] == 0 + + def test_a_demand_survives_being_asked_twice_without_being_counted_twice(self): + """The mirror: room throughout, so the recorded counter must not move.""" + bm = make_block_manager(demand_config(), state_runtime=BIG_IMAGE_RUNTIME) + run_prompt_on_the_ladder(bm, stateful_seq(PROMPT)) + + second = stateful_seq(PROMPT) + for _ in range(5): + assert bm.can_allocate(second) >= 0 + + assert second.checkpoint_demand_pos, "the demand was lost" + funnel = bm.checkpoint_funnel() + assert funnel["demands_recorded"] == 1, "counted per attempt" + assert funnel["demands_declined_no_room"] == 0 + def test_reuse_another_class_declines_is_not_charged_to_the_ladder(self): """The counterfactual keeps every other gate applied. @@ -1673,3 +1874,105 @@ def test_a_demand_is_out_of_generation_s_reach(self): second = stateful_seq(PROMPT) bm.allocate(second, bm.can_allocate(second)) assert second.checkpoint_demand_pos < second.num_prompt_tokens + + +class TestTheCacheCannotStarveLiveKv: + """Why no floor is held back for live KV. + + `_fresh_block` raises when the pool is dry and nothing is evictable, and + the checkpoint cache shares that pool. These pin the three facts that keep + the cache from ever taking it there, so that a future reader looking for a + reserve finds the argument instead of re-inventing one. + """ + + def _pool_of_three_with_one_checkpoint(self, pin: bool): + """A pool holding exactly one image, optionally being read. + + Three units, one checkpoint, nothing spare -- the tightest state the + cache can put the pool in. + """ + bm = make_block_manager( + paged_copy_config(num_kvcache_blocks=3, state_checkpoint_interval_tokens=0), + state_runtime=PAGED_COPY_RUNTIME, + ) + checkpoints = bm.paged_state_checkpoints + assert checkpoints.store.units_per_checkpoint == 3 + assert checkpoints.store.begin_store(33, src_slot=0) is not None + checkpoints.store.complete_inflight() + if pin: + assert checkpoints.begin_restore(33, dst_slot=1) + assert bm.kv.num_free == 0, "the pool is meant to have nothing spare" + return bm, checkpoints.store + + def test_a_ready_unpinned_checkpoint_is_available_to_live_kv(self): + """The cache's size is not the variable: a spendable image is free space. + + `has_available_units` counts it and `ensure_free_units` spends it, so + holding checkpoints costs live KV nothing and there is nothing for a + floor to ration. + """ + bm, store = self._pool_of_three_with_one_checkpoint(pin=False) + assert store.has_available_units(3), "the image was not counted as free space" + assert not store.has_available_units(4), "more was counted than exists" + + seq = stateful_seq(list(range(BLOCK))) + assert bm.can_allocate(seq) == 0, "a spendable image was not counted" + + bm.allocate(seq, 0) + assert store.lookup(33) < 0, "it was counted but could not be spent" + + def test_a_pinned_cache_refuses_an_admission_rather_than_raising(self): + """The reachable outcome under contention, and the one that is not. + + A restore pin is the one thing that makes an image unspendable while + allocation is running. Even with the whole cache pinned and the free + list empty, the gate answers no and the request waits for the pass that + releases the pin -- `_fresh_block` is never reached. + """ + bm, store = self._pool_of_three_with_one_checkpoint(pin=True) + assert not store.has_available_units(1), "the cache is meant to be pinned" + + seq = stateful_seq(list(range(BLOCK))) + + assert bm.can_allocate(seq) < 0, "the gate admitted a seq it cannot serve" + + def test_bypassing_the_gate_reaches_the_raise(self): + """The sibling that gives the test above its meaning. + + Without this one, `can_allocate` returning -1 would be indistinguishable + from a scenario that was never tight enough to matter. + """ + bm, _ = self._pool_of_three_with_one_checkpoint(pin=True) + seq = stateful_seq(list(range(BLOCK))) + + with pytest.raises(AssertionError, match="No PAGE unit"): + bm.allocate(seq, 0) + + def test_a_pass_releases_the_previous_pins_before_it_allocates(self): + """Why the decode loop never sees a pin. + + Pins live one pass: `schedule` releases the previous batch's before it + admits anything. Observable here because the admission below is only + possible once the pinned image becomes spendable again. + """ + scheduler = make_scheduler( + paged_copy_config(num_kvcache_blocks=3, state_checkpoint_interval_tokens=0), + state_runtime=PAGED_COPY_RUNTIME, + ) + # The same pinned pool as the tests above, built inside a scheduler, + # with the batch that reads the restore gone out -- which is what the + # pin is waiting on. + checkpoints = scheduler.block_manager.paged_state_checkpoints + assert checkpoints.store.begin_store(33, src_slot=0) is not None + checkpoints.store.complete_inflight() + assert checkpoints.begin_restore(33, dst_slot=1) + scheduler.block_manager.take_state_maintenance_ops() + assert not checkpoints.store.has_available_units(1) + assert ( + scheduler.block_manager.can_allocate(stateful_seq(list(range(BLOCK)))) < 0 + ) + + scheduler.add(stateful_seq(list(range(BLOCK)))) + _, scheduled = scheduler.schedule() + + assert scheduled, "the pass allocated before releasing the previous pin" diff --git a/tests/test_state_transfer.py b/tests/test_state_transfer.py index 7f228d7baa..26b2187846 100644 --- a/tests/test_state_transfer.py +++ b/tests/test_state_transfer.py @@ -13,7 +13,7 @@ StateTransfer, ) -COPY_SPEC = PagedStateCheckpointSpec(10, 25, "layout-v2") +COPY_SPEC = PagedStateCheckpointSpec(10, 25, "layout-v2", image_bytes=25) def test_wire_round_trip_keeps_the_complete_capability(): diff --git a/tests/test_v4_checkpoint_slot_copy.py b/tests/test_v4_checkpoint_slot_copy.py new file mode 100644 index 0000000000..9b5ac0d420 --- /dev/null +++ b/tests/test_v4_checkpoint_slot_copy.py @@ -0,0 +1,667 @@ +# SPDX-License-Identifier: MIT + +"""Which bytes of a DeepSeek-V4 Active Slot a checkpoint carries, and that a +store/restore round trip moves exactly those and nothing else. + +`checkpoint_ranges_for` (tested in `test_state_arena.py`) says which bytes of the +compressor *arena* are live. This file covers the step after it: composing +those with the sliding-window rows that share the slot, turning the result into +byte segments at a slot's real address, and round-tripping them through the +copy planner. + +The builder is exercised through unbound methods on a stub rather than a real +`DeepseekV4AttentionMetadataBuilder`, which would want a ModelRunner, a model +and a GPU. What the stub supplies is exactly what these methods read, so the +arithmetic under test is the shipped arithmetic. +""" + +from __future__ import annotations + +import ctypes +from dataclasses import replace +from itertools import pairwise +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +# Every class below reads unbound methods off the builder, and that module +# does `from aiter import dtypes` at load, which the non-GPU CI runner cannot +# satisfy. Asked of the module actually needed rather than of `aiter`: there +# `aiter` resolves as a namespace package -- the failure reads "cannot import +# name 'dtypes' from 'aiter' (unknown location)", not "no module named" -- so +# a guard on `aiter` can succeed and leave the real import to fail anyway. +# +# `exc_type` is not optional here. The module *is* found, so bare +# `importorskip` treats an ImportError out of it as the caller's mistake: a +# deprecation warning on pytest 9.0 and an error from 9.1, which is what CI +# runs. Naming the type is what the warning itself prescribes, and it keeps +# the skip narrow -- anything that is not an ImportError still fails. +Builder = pytest.importorskip( + "atom.model_ops.attentions.deepseek_v4_attn", + reason="the V4 builder's module imports aiter at load", + exc_type=ImportError, +).DeepseekV4AttentionMetadataBuilder + +from atom.model_ops.attentions.paged_state_copy import plan_segmented_copy +from atom.model_ops.attentions.state_arena import ( + StateField, + checkpoint_ranges_for, + entry_bytes_for, + field_extents, +) +from atom.model_ops.attentions.v4_pool_geometry import CSA_RATIO, HCA_RATIO + +NEG_INF = float("-inf") +ROW_BYTES = 64 +SLOTS = 3 + +# DeepSeek-V4's field list in miniature: two CSA families, one HCA family that +# a checkpoint owes nothing, and a draft window that is a sliding window and so +# stays whole. Same order as `_state_fields`, which is the order the bytes are +# seen in. +FIELDS = [ + StateField("csa_main_kv", 2, (4, 8), torch.float32), + StateField("csa_main_score", 2, (4, 8), torch.float32, NEG_INF), + StateField("hca_main_kv", 2, (16, 8), torch.float32, in_checkpoint=False), + StateField( + "hca_main_score", 2, (16, 8), torch.float32, NEG_INF, in_checkpoint=False + ), + StateField("state_window", 1, (6, 8), torch.float32), +] +ARENA_BYTES = entry_bytes_for(FIELDS) +ARENA_ROWS = -(-ARENA_BYTES // ROW_BYTES) +ENTRY_ROWS = 5 # the windows that share the slot, in row-space rows +# Rows of that entry a window position actually reaches. Row 2 is interleave +# padding: the real geometry has some, and an image that carried it would pass +# every test below while still being bigger than it has to be. +ENTRY_ROW_RUNS = [(0, 2), (3, 2)] +ENTRY_LIVE_ROWS = sum(count for _, count in ENTRY_ROW_RUNS) +ENTRY_PAD_ROWS = [2] +SLOT_ROWS = ARENA_ROWS + ENTRY_ROWS + 1 # +1 so there is tail padding to drop +SLOT_BYTES = SLOT_ROWS * ROW_BYTES + + +class _Geo: + """Only what `_checkpoint_slot_ranges` and `_slot_views` read.""" + + arena_rows = ARENA_ROWS + entry_rows = ENTRY_ROWS + slot_rows = SLOT_ROWS + + def slot_span(self, physical: int) -> tuple[int, int]: + return physical * SLOT_ROWS, (physical + 1) * SLOT_ROWS + + def physical_slot(self, group: int) -> int: + return group + + def entry_row_runs(self) -> list[tuple[int, int]]: + return list(ENTRY_ROW_RUNS) + + +class _StubBuilder: + """Stands in for the parts of the builder these two methods touch.""" + + # The real methods, not copies of them: between them they read nothing the + # stub does not supply, and a reimplementation here would stop tracking the + # ones it stands in for. + _assert_ratios_divide_the_alignment = Builder._assert_ratios_divide_the_alignment + _checkpoint_slot_ranges = Builder._checkpoint_slot_ranges + _checkpoint_slot_bases = Builder._checkpoint_slot_bases + _checkpoint_segment_sizes = Builder._checkpoint_segment_sizes + checkpoint_image_bytes = Builder.checkpoint_image_bytes + + def __init__(self, plane: torch.Tensor, alignment: int = 256): + self.pool_geometry = _Geo() + self._arena_planes = [FIELDS] + self._checkpoint_range_cache = None + self._checkpoint_slot_base_cache = None + # What the ladder actually rounds a checkpoint to, which is where the + # guard has to look: `kv_cache_block_size * dcp_world_size`. + self.model_runner = SimpleNamespace( + config=SimpleNamespace( + kv_cache_block_size=alignment, decode_context_parallel_size=1 + ) + ) + # The ratios the guard reads come from the model, not from this file's + # constants -- two dense layers and then the CSA/HCA alternation a V4 + # config declares. + self.compress_ratios = [0, 0, CSA_RATIO, HCA_RATIO, CSA_RATIO, HCA_RATIO] + self._plane = plane + + def _plane_row_widths(self): + return [ROW_BYTES] + + def _slot_views(self): + geo = self.pool_geometry + return [ + [self._plane[slice(*geo.slot_span(geo.physical_slot(g)))]] + for g in range(SLOTS) + ] + + +@pytest.fixture +def plane(): + # uint8 so a row is ROW_BYTES elements and byte offsets are element offsets. + return torch.zeros(SLOTS * SLOT_ROWS, ROW_BYTES, dtype=torch.uint8) + + +def field_spans() -> dict[str, tuple[int, int]]: + """Each field's own byte range inside an entry, padding excluded. + + Spelled out here rather than read back from the arena so the round trip is + checked against the layout it is supposed to have, not against the same + arithmetic it is exercising. + """ + spans = {} + offset = 0 + for field in FIELDS: + align = max(256, field.align) + offset = -(-offset // align) * align + spans[field.name] = (offset, offset + field.bytes_per_entry) + offset += field.bytes_per_entry + return spans + + +def dead_arena_span() -> tuple[int, int]: + """Byte range spanned by the two HCA fields, which are dead together.""" + spans = field_spans() + return spans["hca_main_kv"][0], spans["hca_main_score"][1] + + +def store_and_restore(stub, image_ptrs, image_sizes, live_bytes, src, dst): + """Scatter group `src` into the image, gather it back into group `dst`. + + Host `memmove` rather than the copy kernel: what is under test is which + bytes the descriptor names, and running it on the CPU keeps the test off + a GPU. The same plan serves both directions, as it does in production. + """ + ranges = Builder._checkpoint_slot_ranges(stub) + plan = plan_segmented_copy( + [nbytes for plane in ranges for _, nbytes in plane], image_sizes, live_bytes + ) + slot_bases = Builder._checkpoint_slot_bases(stub) + for group, forward in ((src, True), (dst, False)): + descriptor = np.empty((plan.num_spans, 3), dtype=np.int64) + plan.write_descriptor( + descriptor, + slot_bases[group][None], + np.asarray(image_ptrs, dtype=np.int64)[None], + forward=forward, + ) + for source, destination, nbytes in descriptor: + ctypes.memmove(int(destination), int(source), int(nbytes)) + + +class TestCheckpointSlotRanges: + def test_the_ranges_skip_the_dead_fields_and_the_slot_tail(self, plane): + (ranges,) = Builder._checkpoint_slot_ranges(_StubBuilder(plane)) + dead_start, dead_end = dead_arena_span() + + covered = set() + for start, nbytes in ranges: + covered |= set(range(start, start + nbytes)) + + assert not covered & set(range(dead_start, dead_end)), "HCA is carried" + # A window is a sliding window, so every row a position reaches is + # carried — and the interleave padding between them is not. + for start, count in ENTRY_ROW_RUNS: + first = (ARENA_ROWS + start) * ROW_BYTES + assert set(range(first, first + count * ROW_BYTES)) <= covered + for row in ENTRY_PAD_ROWS: + first = (ARENA_ROWS + row) * ROW_BYTES + assert not covered & set( + range(first, first + ROW_BYTES) + ), "interleave padding is carried" + # Neither the padding the arena is rounded up by nor the slot's own + # tail alignment belongs to anyone. + assert not covered & set( + range((ARENA_ROWS + ENTRY_ROWS) * ROW_BYTES, SLOT_BYTES) + ) + + def test_the_ranges_are_ordered_disjoint_and_inside_the_slot(self, plane): + (ranges,) = Builder._checkpoint_slot_ranges(_StubBuilder(plane)) + + assert ranges == sorted(ranges) + for (a_start, a_bytes), (b_start, _) in pairwise(ranges): + assert a_start + a_bytes <= b_start + for start, nbytes in ranges: + assert 0 <= start and start + nbytes <= SLOT_BYTES + + def test_the_image_size_is_the_arena_live_bytes_plus_the_windows(self, plane): + stub = _StubBuilder(plane) + + assert Builder.checkpoint_image_bytes(stub) == ( + sum(n for _, n in checkpoint_ranges_for(FIELDS)) + + ENTRY_LIVE_ROWS * ROW_BYTES + ) + assert Builder.checkpoint_image_bytes(stub) < SLOT_BYTES + + @pytest.mark.parametrize( + "block_size, dcp", + [ + (64, 1), # under HCA's ratio outright + (96, 1), # a multiple of neither + (32, 2), # 64 tokens of alignment: dcp does not rescue it + ], + ) + def test_an_alignment_a_ratio_does_not_divide_is_refused( + self, plane, block_size, dcp + ): + """HCA owes nothing only while a checkpoint lands on a pool boundary. + + The guard reads the alignment the ladder uses, `kv_cache_block_size * + decode_context_parallel_size`, against the ratios the model declares. + """ + stub = _StubBuilder(plane) + stub.model_runner.config.kv_cache_block_size = block_size + stub.model_runner.config.decode_context_parallel_size = dcp + + with pytest.raises(ValueError, match="not a multiple of compress"): + Builder.checkpoint_image_bytes(stub) + + @pytest.mark.parametrize("ratio", [3, 512, 96]) + def test_a_model_ratio_the_shipped_alignment_misses_is_refused(self, ratio, plane): + """The reachable way to break the premise, at the block size that ships. + + `config.py` forces `kv_cache_block_size` to 256 for every `DeepseekV4*` + architecture, so a guard written against this file's `CSA_RATIO` / + `HCA_RATIO` could never fire -- 256 is a multiple of both, and no + configuration changes that. What a variant *can* change is + `hf_config.compress_ratios`, and a stride 256 does not divide leaves + HCA's first pool straddling the boundary, which is the silent-accuracy + failure the guard exists for. + """ + stub = _StubBuilder(plane) + assert stub.model_runner.config.kv_cache_block_size == 256 + stub.compress_ratios = [0, CSA_RATIO, ratio] + + with pytest.raises(ValueError, match="not a multiple of compress"): + Builder.checkpoint_image_bytes(stub) + + @pytest.mark.parametrize("block_size, dcp", [(0, 1), (256, 0)]) + def test_an_alignment_of_zero_is_refused_as_itself(self, plane, block_size, dcp): + """Every ratio divides zero, so the ratio check cannot speak for this. + + Folded into that check it would find nothing wrong and raise anyway, + reporting an alignment that "is not a multiple of compress ratios + `[]`" -- an accusation with no ratio in it, about a pool geometry that + simply cannot exist. Its own check, so it can say that instead. + """ + stub = _StubBuilder(plane) + stub.model_runner.config.kv_cache_block_size = block_size + stub.model_runner.config.decode_context_parallel_size = dcp + + with pytest.raises(ValueError, match="not a pool geometry that can exist"): + Builder.checkpoint_image_bytes(stub) + + def test_dcp_can_make_an_alignment_legal(self, plane): + """The product is what matters, so a small block can still divide.""" + stub = _StubBuilder(plane) + stub.model_runner.config.kv_cache_block_size = 64 + stub.model_runner.config.decode_context_parallel_size = 2 # 128 tokens + + assert Builder.checkpoint_image_bytes(stub) > 0 + + +class TestRoundTrip: + """Store slot 0 into PAGE units, gather it back into slot 1.""" + + def test_the_live_bytes_survive_and_the_dead_bytes_are_not_touched(self, plane): + stub = _StubBuilder(plane) + live_bytes = Builder.checkpoint_image_bytes(stub) + + # Two distinguishable fills, so a byte that failed to move and a byte + # that moved when it should not have both show up. + plane[0 * SLOT_ROWS : 1 * SLOT_ROWS] = 0xA5 + plane[1 * SLOT_ROWS : 2 * SLOT_ROWS] = 0x3C + before = plane[1 * SLOT_ROWS : 2 * SLOT_ROWS].clone() + + # A checkpoint image: arbitrary units, deliberately not contiguous and + # not in address order, which is the whole point of PAGE backing. + units = torch.full((4, live_bytes // 2), 0x11, dtype=torch.uint8) + chosen = (2, 0, 3) + store_and_restore( + stub, + [units[i].data_ptr() for i in chosen], + [units.shape[1]] * len(chosen), + live_bytes, + src=0, + dst=1, + ) + + after = plane[1 * SLOT_ROWS : 2 * SLOT_ROWS].reshape(-1) + source = plane[0 * SLOT_ROWS : 1 * SLOT_ROWS].reshape(-1) + untouched = before.reshape(-1) + + # Derive what must have moved from the layout, NOT from + # `_checkpoint_slot_ranges`. Checking the ranges against themselves passes + # for any self-consistent mistake — including a range shifted the same + # way on both sides, which is exactly the bug worth catching here. + dead_start, dead_end = dead_arena_span() + spans = field_spans() + entry_runs = [ + ( + (ARENA_ROWS + start) * ROW_BYTES, + (ARENA_ROWS + start + count) * ROW_BYTES, + ) + for start, count in ENTRY_ROW_RUNS + ] + for start, end in ( + spans["csa_main_kv"], + spans["csa_main_score"], + spans["state_window"], # a draft window, behind the dead fields + *entry_runs, + ): + assert torch.equal( + after[start:end], source[start:end] + ), f"bytes [{start}, {end}) did not survive the round trip" + + # The dead fields must still hold slot 1's own fill: carrying them + # would be waste, writing them from anywhere else would be a bug. + assert torch.equal( + after[dead_start:dead_end], untouched[dead_start:dead_end] + ), "the dead HCA fields were overwritten" + tail = (ARENA_ROWS + ENTRY_ROWS) * ROW_BYTES + assert torch.equal( + after[tail:], untouched[tail:] + ), "the slot's tail padding was overwritten" + for row in ENTRY_PAD_ROWS: + first = (ARENA_ROWS + row) * ROW_BYTES + assert torch.equal( + after[first : first + ROW_BYTES], + untouched[first : first + ROW_BYTES], + ), "the entry's interleave padding was overwritten" + + def test_a_restore_does_not_read_past_the_image(self, plane): + """`total_bytes` is the image, so the tail of the last unit is spare.""" + stub = _StubBuilder(plane) + live_bytes = Builder.checkpoint_image_bytes(stub) + units = torch.full((live_bytes + 4096,), 0x77, dtype=torch.uint8) + ranges = Builder._checkpoint_slot_ranges(stub) + + plan = plan_segmented_copy( + [nbytes for p in ranges for _, nbytes in p], [units.numel()], live_bytes + ) + descriptor = np.empty((plan.num_spans, 3), dtype=np.int64) + plan.write_descriptor( + descriptor, + Builder._checkpoint_slot_bases(stub)[2][None], + np.array([[units.data_ptr()]], dtype=np.int64), + forward=False, + ) + + assert descriptor[:, 2].sum() == live_bytes + assert (descriptor[:, 0] + descriptor[:, 2]).max() - units.data_ptr() == ( + live_bytes + ) + + +class TestTheBuilderDeclaresWhatItDrops: + """`_state_fields` is where the rule actually lives. + + Everything above uses its own field list, so none of it says anything + about what DeepSeek-V4 itself declares — flipping `in_checkpoint` back on + in the builder left every one of those tests green. This is the one that + notices. + """ + + @staticmethod + def builder_stub(): + class _Stub: + _state_dtype = torch.float32 + csa_layers = (2, 4, 6) + hca_layers = (3, 5) + csa_main_state_shape = (13, 1024) + csa_idx_state_shape = (13, 256) + hca_main_state_shape = (133, 512) + head_dim = 512 + rope_head_dim = 64 + index_head_dim = 128 + block_size = 256 + win_with_spec = 133 + compress_ratios = (0, 0, 4, 128, 4, 128, 4, -1) + _kv_fp8 = False + _indexer_fp4 = False + _field_window_dtype = torch.bfloat16 + _field_window_layers = (43,) + _window_field_row_bytes = Builder._window_field_row_bytes + _state_fields = Builder._state_fields + _geometry_ratios = Builder._geometry_ratios + state_transfer = Builder.state_transfer + + return _Stub() + + @classmethod + def build_fields(cls) -> list[StateField]: + return Builder._state_fields(cls.builder_stub()) + + def test_hca_is_the_only_thing_dropped(self): + dropped = {f.name for f in self.build_fields() if not f.in_checkpoint} + + assert dropped == {"hca_main_kv", "hca_main_score"}, ( + "HCA pools [P, P+128) with no overlap, so a resumer writes every " + "row it reads; nothing else here is known to be dead at a boundary" + ) + + def test_the_dropped_fields_are_outside_every_range(self): + fields = self.build_fields() + carried = set() + for start, nbytes in checkpoint_ranges_for(fields): + carried |= set(range(start, start + nbytes)) + + for field, start, end in field_extents(fields): + overlap = carried & set(range(start, end)) + if field.in_checkpoint: + assert overlap, f"{field.name} is carried but has no range" + else: + assert not overlap, f"{field.name} is dropped but has one" + + def test_what_is_dropped_is_versioned_into_the_layout_id(self): + """Two workers disagreeing on the rule read one image two ways. + + Asks `state_transfer` for the id rather than re-deriving the rule + beside it: re-deriving asserts the rule twice and the fence not at + all, so dropping the `nocopy=` segment would leave it green. + """ + stub = self.builder_stub() + + layout_id = Builder.state_transfer(stub).paged_layout_id + + assert ":nocopy=hca_main_kv,hca_main_score" in layout_id + # The packing rule shares the id, and both are fenced by the version. + assert ":entry=packed" in layout_id + assert layout_id.startswith("dsv4-paged-state-v3:") + + def test_the_layout_id_moves_when_the_rule_does(self): + """The fence has to notice a field changing sides, not just exist.""" + stub = self.builder_stub() + before = Builder.state_transfer(stub).paged_layout_id + + stub._state_fields = lambda: [ + replace(f, in_checkpoint=True) for f in Builder._state_fields(stub) + ] + after = Builder.state_transfer(stub).paged_layout_id + + assert after != before + assert ":nocopy=:" in after, "the id moved for some other reason" + + +class TestPageUnitAddressesAreArithmetic: + """The addresses `_page_unit_regions` computes are the ones slicing gave. + + Replacing 22 throwaway tensor views per unit with three multiplications is + only safe if it lands on the same bytes, so this asks both and compares. + The old expression is written out here rather than kept in production: it + is the oracle, not a fallback. + """ + + N_CSA = 3 + NUM_BLOCKS = 5 + ENVELOPE_ROWS = 7 + ROW_BYTES = 16 + IDX_ROWS = 4 + IDX_ROW_BYTES = 8 + + def build(self): + plane = torch.zeros( + self.NUM_BLOCKS * self.ENVELOPE_ROWS, self.ROW_BYTES, dtype=torch.uint8 + ) + idx = torch.zeros( + self.N_CSA, + self.NUM_BLOCKS, + self.IDX_ROWS, + self.IDX_ROW_BYTES, + dtype=torch.uint8, + ) + + class _Runner: + v4_csa_idx_kv = idx + v4_kv_plane = plane + v4_kv_plane_rope = None + + class _Geo: + envelope_rows = self.ENVELOPE_ROWS + + class _Stub: + _page_unit_regions = Builder._page_unit_regions + _page_unit_bases = Builder._page_unit_bases + _page_unit_stream_sizes = Builder._page_unit_stream_sizes + _kv_planes = Builder._kv_planes + model_runner = _Runner() + pool_geometry = _Geo() + csa_layers = tuple(range(self.N_CSA)) + _indexer_fp4 = False + _page_unit_region_cache = None + _page_unit_region_owners = () + + def _plane_row_widths(self): + return [16] + + return _Stub(), plane, idx + + def sliced(self, plane, idx, block_id): + """What a PAGE unit's regions were before, view by view.""" + start = block_id * self.ENVELOPE_ROWS + views = [plane[start : start + self.ENVELOPE_ROWS]] + views += [idx[layer, block_id] for layer in range(self.N_CSA)] + return [(int(v.data_ptr()), v.numel() * v.element_size()) for v in views] + + def test_every_region_lands_where_a_slice_would_have(self): + stub, plane, idx = self.build() + _, sizes = Builder._page_unit_regions(stub) + + for block_id in range(self.NUM_BLOCKS): + (bases,) = Builder._page_unit_bases(stub, [[block_id]]) + got = list(zip(map(int, bases), map(int, sizes), strict=True)) + assert got == self.sliced(plane, idx, block_id), f"block {block_id}" + + def test_an_image_is_its_units_in_order_addresses_and_sizes_together(self): + """The plan is cut by one of these and addressed through the other. + + Unit major, region minor in both. Were the two to disagree, a plan cut + against one order and addressed through the other would put whole + regions in the wrong unit — a silently wrong checkpoint, not a crash. + """ + stub, plane, idx = self.build() + chosen = [3, 0, 4] + + (bases,) = Builder._page_unit_bases(stub, [chosen]) + sizes = Builder._page_unit_stream_sizes(stub, len(chosen)) + + assert list(zip(map(int, bases), map(int, sizes), strict=True)) == [ + pair for block in chosen for pair in self.sliced(plane, idx, block) + ] + + def test_the_regions_are_worked_out_once(self): + stub, _, _ = self.build() + + first = Builder._page_unit_regions(stub) + assert Builder._page_unit_regions(stub) is first + + def test_a_non_contiguous_pool_is_refused_not_mis_addressed(self): + """Affine addressing assumes the layout; say so rather than guess.""" + stub, _, idx = self.build() + stub.model_runner.v4_csa_idx_kv = idx.transpose(0, 1) + + with pytest.raises(RuntimeError, match="contiguous"): + Builder._page_unit_regions(stub) + + +class TestWarmup: + """That the checkpoint copy path is paid for before a request can pay it.""" + + def test_a_backend_without_paged_checkpoints_warms_nothing(self): + """The hook is on the base builder, so every backend answers it.""" + from atom.model_ops.attentions.backends import AttentionMetadataBuilder + + # No-op by contract: a backend that never copies has nothing to warm, + # and ModelRunner calls this unconditionally. + assert AttentionMetadataBuilder.warmup_per_req_cache(object()) is None + + def test_the_runner_warms_the_builder_once_the_pools_are_reachable(self): + """After the setattr loop, not before: the builder reads them off it. + + Guarded here rather than left to review because nothing else calls + `warmup_per_req_cache` -- a runner that stopped would restore the + first-request JIT silently, with every test still green. + """ + import inspect + + from atom.model_engine import model_runner + + src = inspect.getsource(model_runner.ModelRunner.allocate_kv_cache) + install = src.index("setattr(self, name, value)") + warm = src.index("warmup_per_req_cache()") + assert install < warm, "warmed before the pools were installed" + + +class TestPageUnitRegionsValidateTheirOwnAddresses: + """The region cache holds addresses from two allocations, one uncovered. + + `_invalidate_pool_caches` is called from `allocate_per_req_cache`, which + owns the KV planes; the indexer pools come from `allocate_kv_cache_tensors` + and it does not. A cache relying on that hook would hold freed base + pointers after any path that re-runs the first allocation alone, and the + copy kernel would scatter into whatever the allocator handed that range to + -- no fault, silent corruption. So it keys on its own addresses. + """ + + def _stub(self, plane, idx): + stub = _StubBuilder.__new__(_StubBuilder) + stub.pool_geometry = SimpleNamespace(envelope_rows=2) + stub._page_unit_region_cache = None + stub._page_unit_region_owners = () + stub._indexer_fp4 = False + stub.csa_layers = [0] + stub.model_runner = SimpleNamespace(v4_csa_idx_kv=idx) + stub._kv_planes = lambda: [plane] + stub._plane_row_widths = lambda: [ROW_BYTES] + return stub + + def test_a_moved_indexer_pool_is_noticed_without_the_hook(self): + first = torch.zeros(4, ROW_BYTES, dtype=torch.uint8) + idx = torch.zeros(1, 4, 8, dtype=torch.uint8) + stub = self._stub(first, idx) + bases, _ = Builder._page_unit_regions(stub) + assert idx.data_ptr() in set(bases.tolist()) + + # The pool is reallocated by the other allocation; nothing calls the + # invalidator, exactly as today's ordering would allow. + moved = torch.zeros(1, 4, 8, dtype=torch.uint8) + assert moved.data_ptr() != idx.data_ptr(), "the test needs a real move" + stub.model_runner.v4_csa_idx_kv = moved + + again, _ = Builder._page_unit_regions(stub) + + assert moved.data_ptr() in set(again.tolist()), "answered from a freed base" + assert idx.data_ptr() not in set(again.tolist()) + + def test_an_unmoved_pool_is_answered_from_the_cache(self): + plane = torch.zeros(4, ROW_BYTES, dtype=torch.uint8) + stub = self._stub(plane, torch.zeros(1, 4, 8, dtype=torch.uint8)) + + first = Builder._page_unit_regions(stub) + + assert Builder._page_unit_regions(stub) is first, "rebuilt for nothing" diff --git a/tests/test_v4_pool_geometry.py b/tests/test_v4_pool_geometry.py index 1dd2fd342b..a63e0474a5 100644 --- a/tests/test_v4_pool_geometry.py +++ b/tests/test_v4_pool_geometry.py @@ -10,6 +10,8 @@ pairwise rather than comparing against a restatement of the same expression. """ +from itertools import pairwise + import pytest from atom.model_ops.attentions.v4_pool_geometry import ( @@ -17,8 +19,10 @@ CSA_RATIO, DENSE_RATIO, HCA_RATIO, + ClassLayout, UnifiedPoolGeometry, entry_rows_for, + merge_abutting, ring_offset_for, ) @@ -125,6 +129,105 @@ def test_offsets_are_injective_across_layers_and_positions(self): assert row < entry_rows_for(num_layers, stride, ring_slots) +class TestTheRowsAWindowActuallyReaches: + """`entry_row_runs` against the addresses, not against its own algebra. + + A checkpoint image leaves out whatever these runs do not cover, so the + claim under test is exactly "no `(layer, position)` maps there". Getting + it wrong loses window rows silently — a resumer reads stale KV at the far + end of its window, which costs a fraction of a point and crashes nothing. + """ + + @staticmethod + def reached(num_layers, stride, ring_slots): + return { + layer * stride + ring_offset_for(num_layers, stride, pos) + for layer in range(num_layers) + for pos in range(ring_slots) + } + + @staticmethod + def covered(runs): + return {row for start, count in runs for row in range(start, start + count)} + + def test_the_runs_are_exactly_the_reachable_rows(self): + # `ring_slots < stride` is included on purpose: it is the `whole == 0` + # branch, where the construction is nothing but per-layer partials, + # and a window shorter than `block_size // CSA_RATIO` takes it. + for num_layers in (1, 2, 3, 5, 20, 21): + for stride in (1, 2, 3, 7, 64): + for ring_slots in (1, 2, 5, 64, 128, 131, 133): + cls = ClassLayout( + ratio=4, + layers=tuple(range(num_layers)), + block_rows=0, + ring_stride=stride, + ring_slots=ring_slots, + envelope_offset=0, + entry_offset=0, + ) + what = (num_layers, stride, ring_slots) + assert self.covered(cls.entry_row_runs()) == self.reached( + num_layers, stride, ring_slots + ), what + + def test_the_runs_are_ordered_disjoint_and_inside_the_entry(self): + for cls in flash_geometry().classes.values(): + runs = cls.entry_row_runs() + assert runs == sorted(runs) + for (a_start, a_count), (b_start, _) in pairwise(runs): + assert a_start + a_count <= b_start + for start, count in runs: + assert 0 <= start and start + count <= cls.entry_rows + + def test_what_they_leave_out_is_what_the_interleave_costs(self): + for cls in flash_geometry().classes.values(): + packed = sum(count for _, count in cls.entry_row_runs()) + assert packed == cls.num_layers * cls.ring_slots + assert cls.entry_rows - packed >= 0 + + def test_a_window_that_divides_the_stride_leaves_nothing_out(self): + """No partial ring run, so the whole class is one contiguous range.""" + cls = ClassLayout( + ratio=4, + layers=tuple(range(21)), + block_rows=0, + ring_stride=64, + ring_slots=128, + envelope_offset=0, + entry_offset=0, + ) + + assert cls.entry_row_runs() == [(0, 21 * 128)] + + def test_the_pool_offsets_each_class_and_merges_where_they_meet(self): + geo = flash_geometry() + + runs = geo.entry_row_runs() + covered = self.covered(runs) + + assert sum(count for _, count in runs) == sum( + c.num_layers * c.ring_slots for c in geo.classes.values() + ) + assert covered == { + cls.entry_offset + row + for cls in geo.classes.values() + for row in self.reached(cls.num_layers, cls.ring_stride, cls.ring_slots) + } + assert max(row for row in covered) < geo.entry_rows + + def test_every_window_index_lands_in_a_run(self): + """The formula the kernels use, checked against the runs directly.""" + geo = flash_geometry() + covered = self.covered(geo.entry_row_runs()) + + for layer_id, ratio in enumerate(FLASH_RATIOS): + cls = geo.layer_class(layer_id) + for pos in range(geo.ring_slots): + row = cls.entry_offset + cls.ring_row(cls.layer_index(layer_id), pos) + assert row in covered, (layer_id, ratio, pos, row) + + class TestFlashLayout: def test_class_shapes(self): geo = flash_geometry() @@ -518,3 +621,30 @@ def test_the_classes_that_remain_are_laid_out_as_if_it_never_existed(self): assert geo.envelope_rows == without.envelope_rows for ratio in (CSA_RATIO, HCA_RATIO): assert geo.window_params(ratio) == without.window_params(ratio) + + +class TestMergeAbuttingRefusesDisorder: + """Ordering is a requirement of `merge_abutting`, not a property of callers. + + Each run is compared only against the one before it, so an unordered input + merges nothing and reads as legal. Nothing downstream can see the result is + wrong: `checkpoint_image_bytes` over-counts, and both the op validator and + the sizing cross-check compare against that same number. + """ + + def test_an_overlapping_run_is_refused(self): + # Returned [(0, 400), (256, 128)] before: 528 bytes claimed for 400. + with pytest.raises(ValueError, match="ascend and not overlap"): + merge_abutting([(0, 400), (256, 64), (320, 64)]) + + def test_a_descending_run_is_refused(self): + with pytest.raises(ValueError, match="ascend and not overlap"): + merge_abutting([(256, 64), (0, 64)]) + + def test_a_negative_run_is_refused(self): + with pytest.raises(ValueError, match="non-negative"): + merge_abutting([(0, -1)]) + + def test_abutting_and_separated_runs_are_unchanged(self): + assert merge_abutting([(0, 4), (4, 4), (12, 4)]) == [(0, 8), (12, 4)] + assert merge_abutting([]) == []