Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 107 additions & 23 deletions atom/model_engine/block_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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.
Expand Down
2 changes: 2 additions & 0 deletions atom/model_engine/llm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
)
}
Expand Down Expand Up @@ -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 = {
Expand Down
15 changes: 13 additions & 2 deletions atom/model_engine/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading