class OffloadingConnectorScheduler:
"""Implementation of Scheduler side methods"""
def __init__(
self,
spec: OffloadingSpec,
vllm_config: VllmConfig,
kv_cache_config: KVCacheConfig,
):
self.config = SchedulerOffloadConfig.from_spec(
spec, vllm_config, kv_cache_config
)
self.manager: OffloadingManager = spec.get_manager()
self._connector_stats = OffloadingConnectorStats()
full_attention_groups: list[int] = []
sliding_window_groups: list[int] = []
for config_idx, group_config in enumerate(self.config.kv_group_configs):
if group_config.sliding_window_size_in_chunks is None:
full_attention_groups.append(config_idx)
else:
sliding_window_groups.append(config_idx)
# sort sliding window groups by window size in decreasing order
def _sliding_window_sort_key(i: int) -> int:
val = self.config.kv_group_configs[i].sliding_window_size_in_chunks
assert val is not None
return val
sliding_window_groups.sort(key=_sliding_window_sort_key, reverse=True)
# used by _lookup
self._sliding_window_groups: tuple[int, ...] = tuple(sliding_window_groups)
self._lookup_groups = tuple(full_attention_groups) + self._sliding_window_groups
self._mamba_align_size: int | None = resolve_mamba_align_size(
spec, kv_cache_config
)
self._partial_tail_block_size = (
self.config.kv_group_configs[0].tokens_per_block
if self.config.supports_partial_tail
else 0
)
self._cow_source_groups = frozenset(
config.group_idx
for config in self.config.kv_group_configs
if config.requires_cow_source
)
self._req_status: dict[ReqId, RequestOffloadState] = {}
self._current_batch_load_jobs: dict[int, TransferJob] = {}
self._current_batch_jobs_to_flush: set[int] = set()
# GPU block IDs allocated in the current engine step
self._current_batch_allocated_block_ids: set[int] = set()
# if GPU prefix caching is enabled,
# Track loaded chunks to avoid redundant loads.
self._chunks_being_loaded: set[OffloadKey] | None = (
set() if vllm_config.cache_config.enable_prefix_caching else None
)
# Job ID counter shared by loads and stores.
self._job_counter: int = 0
# Threshold value for stale jobs. All job ids >= _stale_job_threshold are
# active jobs.
self._stale_job_threshold: int = 0
self._jobs: dict[int, TransferJobStatus] = {}
# block_id -> pending store job_ids. Used to track jobs that needs
# flushing in case a block is re-allocated by the KV cache manager.
# Populated only for finished requests (running-request blocks are
# protected by their ref_cnt) and for sliding window blocks (which can
# be freed before a request finishes).
self._block_id_to_pending_jobs: dict[int, set[int]] = {}
self._events_tracker = OffloadingEventsTracker(spec.kv_events_config)
def _maybe_observe_lookup_async_delay(
self, req_status: RequestOffloadState
) -> None:
start_time = req_status.deferred_lookup_start_time
if start_time is None:
return
req_status.deferred_lookup_start_time = None
self._connector_stats.observe_histogram(
_ConnectorMetricName.LOOKUP_ASYNC_DELAY,
time.monotonic() - start_time,
)
def _generate_job_id(self) -> int:
job_id = self._job_counter
self._job_counter += 1
return job_id
def _remove_pending_job(self, job_id: int, block_ids: list[int] | None) -> None:
for bid in block_ids or ():
pending = self._block_id_to_pending_jobs[bid]
pending.remove(job_id)
if not pending:
del self._block_id_to_pending_jobs[bid]
def _calc_num_offloadable_tokens(
self, req_status: RequestOffloadState, num_computed_tokens: int
) -> int:
num = min(num_computed_tokens, req_status.req.num_tokens)
max_offload_tokens = req_status.max_offload_tokens
if max_offload_tokens is not None:
num = min(num, max_offload_tokens)
if self.config.offload_prompt_only:
num = min(num, req_status.req.num_prompt_tokens)
return num
def _maximal_prefix_lookup(
self,
keys: Iterable[OffloadKey],
req_context: ReqContext,
req: Request,
group_config: GroupOffloadConfig,
start_chunk_idx: int,
) -> int | None:
"""Return the number of consecutive offloaded chunks from the start,
or None if the backend deferred a lookup."""
hit_count = 0
defer_lookup = False
for local_idx, key in enumerate(keys):
result = self.manager.lookup(key, req_context)
match result:
case LookupResult.HIT:
self._events_tracker.record_lookup(
req,
group_config,
start_chunk_idx + local_idx,
key,
)
hit_count += 1
case LookupResult.HIT_PENDING:
defer_lookup = True
hit_count += 1
case LookupResult.RETRY:
# Don't break: keep scanning to let manager kick off
# async lookups (until a miss is detected).
defer_lookup = True
case LookupResult.MISS:
break
return hit_count if not defer_lookup else None
def _sliding_window_lookup(
self,
keys: Sequence[OffloadKey],
sliding_window_size: int,
req_context: ReqContext,
initial_window_size: int | None = None,
) -> int | None:
"""Return the end index (in `keys`) of the last run of
`sliding_window_size` consecutive hits, scanning from the end.
The first run may need a larger window for a partial rightmost chunk.
Returns 0 on miss, None if the backend deferred a lookup."""
defer_lookup = False
pending_in_window = False
consecutive_hits = 0
required_window = initial_window_size or sliding_window_size
for idx in range(len(keys) - 1, -1, -1):
match self.manager.lookup(keys[idx], req_context):
case LookupResult.HIT:
consecutive_hits += 1
case LookupResult.HIT_PENDING:
# Block is in cache, just not readable yet — counts
# as hit for the consecutive streak. Don't break:
# keep scanning to let manager kick off async lookups.
pending_in_window = True
consecutive_hits += 1
case LookupResult.RETRY:
# Block location uncertain — does not count as hit.
# Don't break: keep scanning to let manager kick off
# async lookups.
defer_lookup = True
consecutive_hits = 0
pending_in_window = False
required_window = sliding_window_size
case LookupResult.MISS:
consecutive_hits = 0
# This gap rules out the incomplete window to its right.
pending_in_window = False
required_window = sliding_window_size
if consecutive_hits == required_window:
return (
None if defer_lookup or pending_in_window else idx + required_window
)
return None if defer_lookup or pending_in_window else consecutive_hits
def _lookup_complete_chunks(
self,
req_status: RequestOffloadState,
max_num_new_tokens: int | None = None,
) -> int | None:
"""Find how many tokens beyond num_locally_computed_tokens can be loaded.
Iterates full-attention groups first (prefix lookup), then sliding-window
groups (suffix lookup). Each group may tighten max_hit_size_tokens, which
can invalidate an earlier group's result, so the loop re-runs when that
happens until num_hit_tokens converges.
"""
num_computed_tokens = req_status.num_locally_computed_tokens
max_hit_size_tokens: int = req_status.req.num_tokens
if max_num_new_tokens is not None:
max_hit_size_tokens = min(
max_hit_size_tokens, num_computed_tokens + max_num_new_tokens
)
if req_status.max_load_tokens is not None:
max_hit_size_tokens = min(
max_hit_size_tokens,
num_computed_tokens + req_status.max_load_tokens,
)
if self._sliding_window_groups:
# the last prompt token has to be recomputed to get the logprobs
# for sliding window attention, we must reduce by 1 to make sure
# we still have a hit after reduction
max_hit_size_tokens -= 1
if self._mamba_align_size is not None:
# Constrain hit-window to the mamba block size.
max_hit_size_tokens = round_down(
max_hit_size_tokens, self._mamba_align_size
)
num_hit_tokens: int = 0
defer_lookup = False
lookup_groups = self._lookup_groups
# Tracks which eagle groups have already popped their volatile trailing chunk
# in the current convergence iteration. Reset when a non-eagle group
# tightens the hit boundary, requiring a fresh pop.
eagle_verified: set[int] = set()
while lookup_groups:
looked_up_sliding_window: bool = False
groups_iter = iter(lookup_groups)
lookup_groups = ()
for group_idx in groups_iter:
group_config: GroupOffloadConfig = self.config.kv_group_configs[
group_idx
]
group_state: RequestGroupState = req_status.group_states[group_idx]
tokens_per_chunk = group_config.tokens_per_chunk
offload_keys = group_state.offload_keys
assert (
len(offload_keys) >= req_status.req.num_tokens // tokens_per_chunk
)
is_eagle_unverified = (
group_config.is_eagle_group and group_idx not in eagle_verified
)
# Constrain to a chunk-aligned boundary for this group.
max_hit_size_tokens = min(
max_hit_size_tokens, len(offload_keys) * tokens_per_chunk
)
if req_status.max_load_tokens is not None:
max_hit_size_tokens = round_down(
max_hit_size_tokens, tokens_per_chunk
)
if max_hit_size_tokens - num_computed_tokens < tokens_per_chunk:
# We can only load less than a chunk, so skip.
return 0
sliding_window_size_in_chunks = (
group_config.sliding_window_size_in_chunks
)
# For eagle groups, query one extra chunk that will be popped.
# Widening applies to every group type: without it, the pop
# below shrinks max_hit_size_tokens past what was queried,
# which can push the confirmed boundary under a coarser
# sibling group's chunk granularity and zero the whole
# request's hit (issue #52735).
query_max = max_hit_size_tokens
if is_eagle_unverified:
query_max = min(
max_hit_size_tokens + tokens_per_chunk,
len(offload_keys) * tokens_per_chunk,
)
num_chunks = min(cdiv(query_max, tokens_per_chunk), len(offload_keys))
start_chunk_idx = num_computed_tokens // tokens_per_chunk
offload_keys = offload_keys[start_chunk_idx:num_chunks]
# end index (in the sliced offload_keys) up to which we
# have backend-confirmed hits
num_hit_chunks: int | None
if sliding_window_size_in_chunks is None:
num_hit_chunks = self._maximal_prefix_lookup(
offload_keys,
req_status.req_context,
req_status.req,
group_config,
start_chunk_idx,
)
else:
required_window = sliding_window_size_in_chunks
if is_eagle_unverified:
required_window += 1
candidate_end = min(
max_hit_size_tokens,
(num_chunks - int(is_eagle_unverified)) * tokens_per_chunk,
)
initial_window = group_config.load_window_size_in_chunks(
candidate_end
)
assert initial_window is not None
num_hit_chunks = self._sliding_window_lookup(
offload_keys,
required_window,
req_status.req_context,
initial_window + int(is_eagle_unverified),
)
if num_hit_chunks == 0:
return 0
if num_hit_chunks is None:
defer_lookup = True
else:
if is_eagle_unverified:
num_hit_chunks -= 1
eagle_verified.add(group_idx)
max_hit_size_tokens = min(
max_hit_size_tokens,
tokens_per_chunk * (start_chunk_idx + num_hit_chunks),
)
new_num_hit_tokens = max_hit_size_tokens - num_computed_tokens
if new_num_hit_tokens < tokens_per_chunk:
# We can only load less than a chunk, so skip.
return 0
if new_num_hit_tokens < num_hit_tokens:
if not group_config.is_eagle_group:
eagle_verified.clear()
if defer_lookup:
# make another iteration on all groups to check
# if we still need to defer lookup
defer_lookup = False
lookup_groups = self._lookup_groups
elif looked_up_sliding_window and not lookup_groups:
# we need another iteration to confirm previously looked up
# sliding window works with the new_num_hit_tokens
lookup_groups = self._sliding_window_groups
looked_up_sliding_window |= sliding_window_size_in_chunks is not None
num_hit_tokens = new_num_hit_tokens
if defer_lookup:
logger.debug(
"Offloading manager delayed request %s as backend requested",
req_status.req.request_id,
)
return None
# Possibly delay the request if any hit chunk is already being loaded.
if self._chunks_being_loaded:
for group_config, group_state in zip(
self.config.kv_group_configs, req_status.group_states
):
tokens_per_chunk = group_config.tokens_per_chunk
sliding_window_size_in_chunks = group_config.load_window_size_in_chunks(
num_computed_tokens + num_hit_tokens
)
offload_keys = group_state.offload_keys
num_chunks = cdiv(
num_computed_tokens + num_hit_tokens, tokens_per_chunk
)
start_chunk_idx = num_computed_tokens // tokens_per_chunk
offload_keys = offload_keys[start_chunk_idx:num_chunks]
if sliding_window_size_in_chunks is not None:
offload_keys = offload_keys[-sliding_window_size_in_chunks:]
if any(key in self._chunks_being_loaded for key in offload_keys):
# Hit chunks are being loaded, so delay the request.
logger.debug(
"Delaying request %s since some of its"
" chunks are already being loaded",
req_status.req.request_id,
)
return None
logger.debug(
"Request %s hit %s offloaded tokens after %s GPU hit tokens",
req_status.req.request_id,
num_hit_tokens,
num_computed_tokens,
)
return num_hit_tokens
def _make_boundary_key(
self,
request: Request,
group_idx: int,
boundary_tokens: int,
req_context: ReqContext,
) -> OffloadKey:
hash_idx = boundary_tokens // self.config.tokens_per_hash - 1
key = make_offload_key(request.block_hashes[hash_idx], group_idx)
req_context.set_offload_key_position(key, boundary_tokens)
return key
def _lookup(
self,
req_status: RequestOffloadState,
max_num_new_tokens: int | None = None,
) -> int | None:
complete_hit = self._lookup_complete_chunks(req_status, max_num_new_tokens)
req_status.partial_tail_boundary = None
if complete_hit is None or not self.config.supports_partial_tail:
return complete_hit
local_tokens = req_status.num_locally_computed_tokens
complete_boundary = local_tokens + complete_hit
tokens_per_hash = self.config.tokens_per_hash
block_end = complete_boundary + self._partial_tail_block_size
max_boundary = min(req_status.req.num_prompt_tokens - 1, block_end - 1)
if max_num_new_tokens is not None:
max_boundary = min(max_boundary, local_tokens + max_num_new_tokens)
if req_status.max_load_tokens is not None:
max_boundary = min(
max_boundary,
local_tokens + req_status.max_load_tokens,
)
max_boundary = round_down(max_boundary, tokens_per_hash)
if max_boundary <= complete_boundary:
return complete_hit
pending = False
for boundary in range(max_boundary, complete_boundary, -tokens_per_hash):
boundary_pending = False
boundary_missed = False
boundary_keys = []
for group_config in self.config.kv_group_configs:
key = self._make_boundary_key(
req_status.req,
group_config.group_idx,
boundary,
req_status.req_context,
)
boundary_keys.append(key)
result = self.manager.lookup(key, req_status.req_context)
if result is LookupResult.MISS:
boundary_missed = True
break
if result in (LookupResult.HIT_PENDING, LookupResult.RETRY):
boundary_pending = True
pending |= boundary_pending
if not boundary_missed and not boundary_pending:
for group_config, key in zip(
self.config.kv_group_configs, boundary_keys
):
self._events_tracker.record_partial_lookup(
req_status.req, group_config, boundary, key
)
req_status.partial_tail_boundary = boundary
return boundary - local_tokens
if pending and complete_hit == 0:
return None
return complete_hit
def on_new_request(self, request: Request) -> None:
"""Called when a new request is added to the scheduler."""
req_context = _create_req_context(request)
offloading_context = self.manager.on_new_request(req_context)
req_status = RequestOffloadState(
config=self.config,
req=request,
req_context=req_context,
offloading_context=offloading_context,
)
self._req_status[request.request_id] = req_status
def get_num_new_matched_tokens(
self,
request: Request,
num_computed_tokens: int,
max_num_new_tokens: int | None = None,
) -> tuple[int | None, bool]:
"""Get number of new tokens that can be loaded beyond the
num_computed_tokens.
Args:
request (Request): the request object.
num_computed_tokens (int): the number of locally
computed tokens for this request
max_num_new_tokens (int | None): cap on the number of tokens that
may be loaded beyond `num_computed_tokens`, if any.
Returns:
A tuple with the following elements:
- The number of tokens that can be loaded beyond what is
already computed.
If None, it means that the connector needs more time to
determine the number of matched tokens, and the scheduler
should query for this request again later.
- `True` if tokens will be loaded asynchronously
(between scheduler steps).
"""
req_status = self._req_status[request.request_id]
for group_state in req_status.group_states:
group_state.block_ids.clear()
if req_status.transfer_jobs:
logger.debug(
"Delaying request %s since it still has in-flight transfers",
request.request_id,
)
return None, False
req_status.update_offload_keys()
req_status.num_locally_computed_tokens = num_computed_tokens
num_hit_tokens: int | None
if request.skip_reading_prefix_cache:
num_hit_tokens = 0
else:
lookup_start = time.monotonic()
num_hit_tokens = self._lookup(req_status, max_num_new_tokens)
self._connector_stats.observe_histogram(
_ConnectorMetricName.LOOKUP_SYNC_DELAY,
time.monotonic() - lookup_start,
)
if num_hit_tokens is None:
if req_status.deferred_lookup_start_time is None:
req_status.deferred_lookup_start_time = lookup_start
else:
self._maybe_observe_lookup_async_delay(req_status)
req_status.update_num_hit_chunks(num_computed_tokens + (num_hit_tokens or 0))
return num_hit_tokens, bool(num_hit_tokens)
def update_state_after_alloc(
self, request: Request, blocks: KVCacheBlocks, num_external_tokens: int
):
if num_external_tokens == 0:
return
req_status = self._req_status[request.request_id]
num_locally_computed_tokens = req_status.num_locally_computed_tokens
num_cached_tokens = num_locally_computed_tokens + num_external_tokens
partial_tail_boundary = req_status.partial_tail_boundary
if partial_tail_boundary is not None:
assert partial_tail_boundary == num_cached_tokens
keys_to_load: list[OffloadKey] = []
dst_block_ids: list[int] = []
# per group
group_sizes: list[int] = []
block_indices: list[int] = []
for group_config, group_state in zip(
self.config.kv_group_configs,
req_status.group_states,
):
group_blocks = blocks.blocks[group_config.group_idx]
self._current_batch_allocated_block_ids.update(
block.block_id for block in group_blocks if block.block_id != 0
)
tokens_per_block = group_config.tokens_per_block
tokens_per_chunk = group_config.tokens_per_chunk
offload_keys = group_state.offload_keys
num_gpu_blocks = cdiv(num_cached_tokens, tokens_per_block)
assert len(group_blocks) >= num_gpu_blocks
# ``load_start_gpu_block_idx``: the index in ``group_blocks`` where the
# load region begins -- a slice bound, not a count of computed blocks.
# Scan from the computed boundary, not 0: sparse groups (Mamba,
# SWA) legitimately hold non-null unhashed blocks below it.
# Skip nulls (sentinel / out-of-retention padding).
first_fresh_gpu_block_idx = cdiv(
num_locally_computed_tokens, tokens_per_block
)
load_start_gpu_block_idx = num_gpu_blocks
for i in range(first_fresh_gpu_block_idx, num_gpu_blocks):
block = group_blocks[i]
if not block.is_null and block.block_hash is None:
load_start_gpu_block_idx = i
break
assert num_locally_computed_tokens % tokens_per_block == 0
num_pending_gpu_blocks = num_gpu_blocks - load_start_gpu_block_idx
if group_config.sliding_window_size_in_chunks is not None:
assert (
num_pending_gpu_blocks
<= group_config.sliding_window_size_in_chunks
* self.config.blocks_per_chunk
+ 1
)
num_chunks = cdiv(num_cached_tokens, tokens_per_chunk)
if num_pending_gpu_blocks:
start_chunk_idx = (
load_start_gpu_block_idx // self.config.blocks_per_chunk
)
end_chunk_idx = num_chunks - (partial_tail_boundary is not None)
assert len(offload_keys) >= end_chunk_idx
keys_to_load.extend(offload_keys[start_chunk_idx:end_chunk_idx])
if partial_tail_boundary is not None:
keys_to_load.append(
self._make_boundary_key(
request,
group_config.group_idx,
partial_tail_boundary,
req_status.req_context,
)
)
dst_block_ids.extend(
block.block_id
for block in group_blocks[load_start_gpu_block_idx:num_gpu_blocks]
)
group_sizes.append(num_pending_gpu_blocks)
block_indices.append(load_start_gpu_block_idx)
# Skip prefix-hit chunks for block-level policy; for
# request-level, next_stored_chunk_idx stays at 0 so all
# chunks (including hits) are offloaded.
if req_status.offloading_context.policy == OffloadPolicy.CHUNK_LEVEL:
group_state.next_stored_chunk_idx = num_chunks
src_spec = self.manager.prepare_load(keys_to_load, req_status.req_context)
dst_spec = GPULoadStoreSpec(
dst_block_ids, group_sizes=group_sizes, block_indices=block_indices
)
load_job_id = self._generate_job_id()
self._current_batch_load_jobs[load_job_id] = TransferJob(
req_id=request.request_id,
src_spec=src_spec,
dst_spec=dst_spec,
)
# a load can only be issued when no other jobs are pending.
assert not req_status.transfer_jobs
req_status.transfer_jobs.add(load_job_id)
self._jobs[load_job_id] = TransferJobStatus(
req_id=request.request_id,
pending_count=self.config.num_workers,
keys=set(keys_to_load),
is_store=False,
)
if self._chunks_being_loaded is not None:
self._chunks_being_loaded.update(keys_to_load)
req_status.partial_tail_boundary = None
def _update_req_states(self, scheduler_output: SchedulerOutput) -> None:
"""Update request states from the Scheduler's output."""
# new_block_ids_end[req_id][i] = end of pre-existing block_ids for
# the i-th sliding window group (before this step's extend).
# Used to detect sliding window blocks that got re-allocated.
new_block_ids_end: dict[str, tuple[int, ...]] = {}
for req_id, new_block_id_groups, preempted in yield_req_data(scheduler_output):
req_status = self._req_status[req_id]
req_status.update_offload_keys()
if preempted:
for group_state in req_status.group_states:
group_state.block_ids.clear()
if new_block_id_groups:
if self._sliding_window_groups:
new_block_ids_end[req_id] = tuple(
len(req_status.group_states[grp_idx].block_ids)
for grp_idx in self._sliding_window_groups
)
req_status.update_block_id_groups(new_block_id_groups)
for group_config in self.config.kv_group_configs:
new_blocks = new_block_id_groups[group_config.group_idx]
for bid in new_blocks:
if bid != 0:
self._current_batch_allocated_block_ids.add(bid)
for copy in scheduler_output.kv_cache_block_copies or ():
self._current_batch_allocated_block_ids.add(copy.dst_block_id)
# Zero out stale block_ids in sliding window groups' pending-store
# positions. Only sliding window groups can have stale entries (blocks
# freed by remove_skipped_blocks then reallocated). Only positions in
# [next_stored_chunk_idx * bsf, end) need checking where end is the
# pre-extend length: earlier positions were already offloaded, later
# ones are fresh allocations from this step.
if self._sliding_window_groups and self._current_batch_allocated_block_ids:
blocks_per_chunk = self.config.blocks_per_chunk
for req_id, req_status in self._req_status.items():
ends = new_block_ids_end.get(req_id)
for i, grp_idx in enumerate(self._sliding_window_groups):
group_state = req_status.group_states[grp_idx]
start = group_state.next_stored_chunk_idx * blocks_per_chunk
end = ends[i] if ends is not None else len(group_state.block_ids)
for j in range(start, end):
if (
group_state.block_ids[j]
in self._current_batch_allocated_block_ids
):
group_state.block_ids[j] = 0
def _build_aligned_boundary_store_jobs(
self, handoffs: dict[str, list[tuple[int, int, int]]]
) -> dict[int, TransferJob]:
store_jobs: dict[int, TransferJob] = {}
num_groups = len(self.config.kv_group_configs)
config_idx_by_group = {
config.group_idx: idx
for idx, config in enumerate(self.config.kv_group_configs)
}
for req_id, entries in handoffs.items():
req_status = self._req_status.get(req_id)
if req_status is None:
continue
req = req_status.req
max_boundary = self._calc_num_offloadable_tokens(req_status, req.num_tokens)
for group_idx, block_id, boundary in entries:
config_idx = config_idx_by_group.get(group_idx)
if config_idx is None:
continue
group_config = self.config.kv_group_configs[config_idx]
if (
block_id == 0
or boundary > max_boundary
or boundary % group_config.tokens_per_chunk != 0
):
continue
key = self._make_boundary_key(
req, group_idx, boundary, req_status.req_context
)
store_output = self.manager.prepare_store([key], req_status.req_context)
if store_output is None:
self._connector_stats.increase_counter(
_ConnectorMetricName.ALLOCATION_FAILURE
)
continue
if not store_output.keys_to_store:
continue
job_id = self._generate_job_id()
req_status.transfer_jobs.add(job_id)
self._block_id_to_pending_jobs.setdefault(block_id, set()).add(job_id)
self._jobs[job_id] = TransferJobStatus(
req_id=req_id,
pending_count=self.config.num_workers,
keys={key},
is_store=True,
fenced_block_ids=[block_id],
)
group_sizes = [0] * num_groups
group_sizes[config_idx] = 1
block_indices = [0] * num_groups
block_indices[config_idx] = (
boundary // group_config.tokens_per_block - 1
)
store_jobs[job_id] = TransferJob(
req_id=req_id,
src_spec=GPULoadStoreSpec(
[block_id],
group_sizes=group_sizes,
block_indices=block_indices,
),
dst_spec=store_output.store_spec,
)
self._events_tracker.record_store(
req,
group_config,
boundary // group_config.tokens_per_chunk - 1,
key,
)
return store_jobs
def _build_partial_tail_store_jobs(
self, scheduler_output: SchedulerOutput
) -> dict[int, TransferJob]:
block_state = scheduler_output.kv_connector_block_state
handoffs = block_state.boundary_state_offloads if block_state else None
if not handoffs:
return {}
store_jobs = self._build_aligned_boundary_store_jobs(handoffs)
if not self.config.supports_partial_tail:
return store_jobs
for req_id, entries in handoffs.items():
entries = [
entry
for entry in entries
if entry[2] % self._partial_tail_block_size != 0
]
if not entries:
continue
req_status = self._req_status.get(req_id)
assert req_status is not None
boundaries = {boundary for _, _, boundary in entries}
assert len(boundaries) == 1
boundary = boundaries.pop()
req = req_status.req
group_states = {
group.group_idx: state
for group, state in zip(
self.config.kv_group_configs, req_status.group_states
)
}
max_boundary = min(
req.num_prompt_tokens,
req_status.max_offload_tokens or req.num_prompt_tokens,
)
assert boundary > 0
assert boundary % self.config.tokens_per_hash == 0
assert boundary <= max_boundary
cow_blocks = {group_idx: block_id for group_idx, block_id, _ in entries}
assert self._cow_source_groups.issubset(cow_blocks)
assert boundary % self._partial_tail_block_size != 0
block_idx = boundary // self._partial_tail_block_size
if any(
group.group_idx not in self._cow_source_groups
and block_idx >= len(group_states[group.group_idx].block_ids)
for group in self.config.kv_group_configs
):
continue
keys = [
self._make_boundary_key(
req, group.group_idx, boundary, req_status.req_context
)
for group in self.config.kv_group_configs
]
block_ids = [
cow_blocks[group.group_idx]
if group.group_idx in self._cow_source_groups
else group_states[group.group_idx].block_ids[block_idx]
for group in self.config.kv_group_configs
]
assert all(block_id != 0 for block_id in block_ids)
store_output = self.manager.prepare_store(keys, req_status.req_context)
if store_output is None:
self._connector_stats.increase_counter(
_ConnectorMetricName.ALLOCATION_FAILURE
)
continue
if not store_output.keys_to_store:
continue
for group_config, key in zip(self.config.kv_group_configs, keys):
if key in store_output.keys_to_store:
self._events_tracker.record_partial_store(
req, group_config, boundary, key
)
group_by_key = {key: idx for idx, key in enumerate(keys)}
accepted_groups = [group_by_key[key] for key in store_output.keys_to_store]
group_sizes = [0] * len(self.config.kv_group_configs)
block_indices = [0] * len(self.config.kv_group_configs)
for group_idx in accepted_groups:
group_sizes[group_idx] = 1
block_indices[group_idx] = block_idx
source_blocks = [block_ids[group_idx] for group_idx in accepted_groups]
job_id = self._generate_job_id()
req_status.transfer_jobs.add(job_id)
for block_id in source_blocks:
self._block_id_to_pending_jobs.setdefault(block_id, set()).add(job_id)
self._jobs[job_id] = TransferJobStatus(
req_id=req_id,
pending_count=self.config.num_workers,
keys=set(store_output.keys_to_store),
is_store=True,
fenced_block_ids=source_blocks,
)
store_jobs[job_id] = TransferJob(
req_id=req_id,
src_spec=GPULoadStoreSpec(
source_blocks,
group_sizes=group_sizes,
block_indices=block_indices,
),
dst_spec=store_output.store_spec,
)
return store_jobs
def _reachable_store_block_mask(
self,
group_config: GroupOffloadConfig,
start_chunk_idx: int,
end_chunk_idx: int,
final_segment_end_chunk_idx: int | None,
reachable_boundaries: tuple[int, ...],
) -> list[bool] | None:
"""Build the block mask for a range of candidate offload chunks."""
blocks_per_chunk = self.config.blocks_per_chunk
kv_cache_spec = group_config.kv_cache_spec
if isinstance(kv_cache_spec, SlidingWindowSpec) and (
kv_cache_spec.extra_retained_tokens
):
# Offload restores all allocated history, including MTP re-prefill
# tokens. Widen only the store mask, not the model's attention window.
kv_cache_spec = replace(
kv_cache_spec,
sliding_window=(
kv_cache_spec.sliding_window + kv_cache_spec.extra_retained_tokens
),
extra_retained_tokens=0,
)
return group_config.manager_cls.reachable_block_mask(
start_block=start_chunk_idx * blocks_per_chunk,
end_block=end_chunk_idx * blocks_per_chunk,
alignment_tokens=self.config.alignment_tokens,
kv_cache_spec=kv_cache_spec,
use_eagle=group_config.is_eagle_group,
retention_interval=self.config.retention_interval,
reachable_boundaries=reachable_boundaries,
dcp_world_size=self.config.dcp_world_size,
final_segment_end_block=(
final_segment_end_chunk_idx * blocks_per_chunk
if final_segment_end_chunk_idx is not None
else None
),
)
def _final_swa_alignment_blocks(
self, group_config: GroupOffloadConfig
) -> int | None:
"""Return representable final-tail alignment for a non-EAGLE SWA group."""
alignment_tokens = self.config.alignment_tokens
if (
alignment_tokens is None
or group_config.is_eagle_group
or not isinstance(group_config.kv_cache_spec, SlidingWindowSpec)
or alignment_tokens % group_config.tokens_per_block != 0
):
return None
return alignment_tokens // group_config.tokens_per_block
def _build_store_jobs(
self,
scheduler_output: SchedulerOutput,
) -> dict[int, TransferJob]:
blocks_per_chunk = self.config.blocks_per_chunk
store_jobs: dict[int, TransferJob] = {}
for req_id in chain(
scheduler_output.num_scheduled_tokens,
scheduler_output.finished_req_ids or (),
):
req_status = self._req_status.get(req_id)
if req_status is None:
continue
req = req_status.req
if req.status is RequestStatus.FINISHED_ABORTED:
num_tokens_after_batch = req.num_computed_tokens
elif req.is_finished():
# Clamp to the GPU prefix cache's commit point. The final sampled
# token's slot is never committed (under spec decode it holds a
# rejected draft's KV), so a block ending there must not be stored.
num_tokens_after_batch = max(req.num_prompt_tokens, req.num_tokens - 1)
else:
num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id]
num_tokens_after_batch = req.num_computed_tokens + num_scheduled_tokens
num_offloadable_tokens = self._calc_num_offloadable_tokens(
req_status, num_tokens_after_batch
)
prompt_offloadable_tokens = self._calc_num_offloadable_tokens(
req_status, req.num_prompt_tokens
)
# Filter out chunks skipped due to sliding window attention / SSM
# or unreachable by the load path's alignment constraints.
new_offload_keys: list[OffloadKey] = []
group_store_ranges: list[tuple[int, int]] = []
reachable_boundaries: tuple[int, ...] = ()
if self.config.retention_interval is not None:
reachable_boundaries = (req.num_prompt_tokens - 1,)
if req.shared_prefix_boundary:
reachable_boundaries += (req.shared_prefix_boundary,)
for group_config, group_state in zip(
self.config.kv_group_configs, req_status.group_states
):
num_chunks = req_status.storable_chunks(
group_config, group_state, num_offloadable_tokens
)
start_chunk_idx = group_state.next_stored_chunk_idx
prompt_horizon_chunks = (
prompt_offloadable_tokens // group_config.tokens_per_chunk
)
final_swa_alignment_blocks = self._final_swa_alignment_blocks(
group_config
)
reconsider_final_swa_tail = (
req.is_finished()
and num_chunks != prompt_horizon_chunks
and self.config.retention_interval is None
and final_swa_alignment_blocks is not None
)
if (
self.config.retention_interval is not None
or group_config.is_eagle_group
):
store_horizon_chunks = None
elif req.is_finished():
store_horizon_chunks = num_chunks
elif num_chunks <= prompt_horizon_chunks:
store_horizon_chunks = prompt_horizon_chunks
else:
# An active decode frontier is not a final request
# boundary. Only fixed alignment tails are reachable.
store_horizon_chunks = None
if reconsider_final_swa_tail:
assert final_swa_alignment_blocks is not None
horizon_blocks = num_chunks * blocks_per_chunk
partial_segment_start_block = (
horizon_blocks - horizon_blocks % final_swa_alignment_blocks
)
partial_segment_start_chunk = (
partial_segment_start_block // blocks_per_chunk
)
start_chunk_idx = min(start_chunk_idx, partial_segment_start_chunk)
group_store_ranges.append((start_chunk_idx, num_chunks))
if group_config.requires_cow_source:
continue
if num_chunks <= start_chunk_idx:
continue
offload_keys = group_state.offload_keys[start_chunk_idx:num_chunks]
# For each chunk, take the last corresponding GPU block. For
# blocks_per_chunk=3 and GPU block IDs 1 5 6 7 2 4 9 3 8,
# this selects GPU blocks 6 4 8.
# A block_id of 0 means either a sliding window / SSM skip
# or a stale entry that was zeroed out — skip it either way.
offload_block_ids = group_state.block_ids[
start_chunk_idx * blocks_per_chunk
+ blocks_per_chunk
- 1 : num_chunks * blocks_per_chunk : blocks_per_chunk
]
assert len(offload_keys) == len(offload_block_ids)
# Use reachable_block_mask to filter unreachable chunks
# (SWA/Mamba sparsity + retention interval).
# reachable_block_mask operates in KV-block coordinates,
# so convert chunk indices to block indices.
block_mask = self._reachable_store_block_mask(
group_config=group_config,
start_chunk_idx=start_chunk_idx,
end_chunk_idx=num_chunks,
final_segment_end_chunk_idx=store_horizon_chunks,
reachable_boundaries=reachable_boundaries,
)
prompt_horizon_block_mask: list[bool] | None = None
if reconsider_final_swa_tail:
prompt_horizon_block_mask = self._reachable_store_block_mask(
group_config=group_config,
start_chunk_idx=start_chunk_idx,
end_chunk_idx=num_chunks,
final_segment_end_chunk_idx=prompt_horizon_chunks,
reachable_boundaries=reachable_boundaries,
)
for key_idx, (offload_key, block_id) in enumerate(
zip(offload_keys, offload_block_ids)
):
if block_id == 0:
continue
abs_chunk_idx = start_chunk_idx + key_idx
if (
reconsider_final_swa_tail
and abs_chunk_idx < group_state.next_stored_chunk_idx
and (
prompt_horizon_block_mask is None
or any(
prompt_horizon_block_mask[
key_idx * blocks_per_chunk + block_idx
]
for block_idx in range(blocks_per_chunk)
)
)
):
continue
# A chunk is reachable if any of its constituent
# blocks is reachable.
if block_mask is not None and not any(
block_mask[key_idx * blocks_per_chunk + b]
for b in range(blocks_per_chunk)
):
continue
new_offload_keys.append(offload_key)
if not new_offload_keys:
req_status.advance_stored_idx(num_offloadable_tokens)
continue
store_output = self.manager.prepare_store(
new_offload_keys, req_status.req_context
)
if store_output is None:
self._connector_stats.increase_counter(
_ConnectorMetricName.ALLOCATION_FAILURE
)
logger.warning("Request %s: cannot store chunks", req_id)
continue
if not store_output.keys_to_store:
req_status.advance_stored_idx(num_offloadable_tokens)
continue
keys_to_store = set(store_output.keys_to_store)
group_sizes: list[int] = []
block_indices: list[int] = []
src_block_ids: list[int] = []
fenced_block_ids: list[int] = []
deferred_fence_block_ids: list[int] = []
for group_config, group_state, store_range in zip(
self.config.kv_group_configs,
req_status.group_states,
group_store_ranges,
):
is_sliding_window = (
group_config.sliding_window_size_in_chunks is not None
)
start_chunk_idx, num_chunks = store_range
block_ids = group_state.block_ids
num_group_blocks = 0
start_gpu_block_idx: int | None = None
for idx, offload_key in enumerate(
group_state.offload_keys[start_chunk_idx:num_chunks]
):
if offload_key not in keys_to_store:
continue
chunk_idx = start_chunk_idx + idx
self._events_tracker.record_store(
req, group_config, chunk_idx, offload_key
)
gpu_block_idx = chunk_idx * blocks_per_chunk
for i in range(blocks_per_chunk):
block_id = block_ids[gpu_block_idx + i]
if block_id == 0:
continue
if start_gpu_block_idx is None:
start_gpu_block_idx = gpu_block_idx + i
src_block_ids.append(block_id)
num_group_blocks += 1
if is_sliding_window:
fenced_block_ids.append(block_id)
else:
deferred_fence_block_ids.append(block_id)
group_sizes.append(num_group_blocks)
block_indices.append(start_gpu_block_idx or 0)
group_state.next_stored_chunk_idx = max(
group_state.next_stored_chunk_idx, num_chunks
)
src_spec = GPULoadStoreSpec(
src_block_ids, group_sizes=group_sizes, block_indices=block_indices
)
dst_spec = store_output.store_spec
job_id = self._generate_job_id()
# a store can only be issued when no load is pending.
if req_status.transfer_jobs:
any_jid = next(iter(req_status.transfer_jobs))
assert self._jobs[any_jid].is_store
req_status.transfer_jobs.add(job_id)
# Watch sliding window blocks as they may get evicted
# before the request finishes
for bid in fenced_block_ids:
self._block_id_to_pending_jobs.setdefault(bid, set()).add(job_id)
# the non-sliding window blocks will be watched only
# when the request finishes
self._jobs[job_id] = TransferJobStatus(
req_id=req_id,
pending_count=self.config.num_workers,
keys=set(keys_to_store),
is_store=True,
deferred_fence_block_ids=deferred_fence_block_ids,
fenced_block_ids=fenced_block_ids or None,
)
store_jobs[job_id] = TransferJob(
req_id=req_id, src_spec=src_spec, dst_spec=dst_spec
)
logger.debug(
"Request %s offloading %s chunks upto %d tokens (job %d)",
req_id,
len(keys_to_store),
num_offloadable_tokens,
job_id,
)
if req.is_finished():
# Register non-sliding-window blocks for flush detection.
for bid in deferred_fence_block_ids:
self._block_id_to_pending_jobs.setdefault(bid, set()).add(job_id)
if bid in self._current_batch_allocated_block_ids:
self._current_batch_jobs_to_flush.add(job_id)
return store_jobs
def build_connector_meta(
self, scheduler_output: SchedulerOutput
) -> KVConnectorMetadata:
self._update_req_states(scheduler_output)
schedule_end_context = ScheduleEndContext(
new_req_ids=[req.req_id for req in scheduler_output.scheduled_new_reqs],
preempted_req_ids=scheduler_output.preempted_req_ids or (),
)
self.manager.on_schedule_end(schedule_end_context)
# Flush jobs for preempted requests.
for req_id in scheduler_output.preempted_req_ids or ():
req_status = self._req_status.get(req_id)
if req_status is None or not req_status.transfer_jobs:
continue
any_jid = next(iter(req_status.transfer_jobs))
assert self._jobs[any_jid].is_store
self._current_batch_jobs_to_flush.update(req_status.transfer_jobs)
# Flush jobs that contain re-allocated blocks.
if (
self._block_id_to_pending_jobs
and not self._block_id_to_pending_jobs.keys().isdisjoint(
self._current_batch_allocated_block_ids
)
):
self._current_batch_jobs_to_flush.update(
jid
for bid in self._current_batch_allocated_block_ids
if bid in self._block_id_to_pending_jobs
for jid in self._block_id_to_pending_jobs[bid]
)
partial_store_jobs = self._build_partial_tail_store_jobs(scheduler_output)
normal_store_jobs = self._build_store_jobs(scheduler_output)
meta = OffloadingConnectorMetadata(
load_jobs=self._current_batch_load_jobs,
store_jobs=partial_store_jobs | normal_store_jobs,
jobs_to_flush=self._current_batch_jobs_to_flush,
)
# All prepare_store calls for finished requests have been issued.
# Signal on_request_finished and clean up state where possible.
for req_id in scheduler_output.finished_req_ids or ():
req_status = self._req_status.get(req_id)
if req_status is None:
continue
req_status.finished_signaled = True
self.manager.on_request_finished(req_status.req_context)
if not req_status.transfer_jobs:
del self._req_status[req_id]
self._current_batch_load_jobs = {}
self._current_batch_jobs_to_flush = set()
self._current_batch_allocated_block_ids = set()
return meta
def has_pending_push_work(self) -> bool:
"""Whether the engine must keep stepping.
While True, build_connector_meta() and update_connector_output()
continue to be called even when no requests are scheduled.
"""
return bool(self._jobs) or self.manager.has_pending_work()
def update_connector_output(self, connector_output: KVConnectorOutput):
"""Update KVConnector state from worker-side connectors output.
Args:
connector_output (KVConnectorOutput): the worker-side
connectors output.
"""
meta = connector_output.kv_connector_worker_meta
if not isinstance(meta, OffloadingWorkerMetadata):
assert meta is None
meta = OffloadingWorkerMetadata()
if not meta.transfer_stats.is_empty():
transfer_stats = OffloadingConnectorStats()
if not meta.transfer_stats.load.is_empty():
transfer_stats.increase_counter(
_TransferMetricName.LOAD_BYTES,
meta.transfer_stats.load.bytes,
)
transfer_stats.increase_counter(
_TransferMetricName.LOAD_TIME,
meta.transfer_stats.load.time,
)
for size in meta.transfer_stats.load.sizes:
transfer_stats.observe_histogram(
_TransferMetricName.LOAD_SIZE, size
)
if not meta.transfer_stats.store.is_empty():
transfer_stats.increase_counter(
_TransferMetricName.STORE_BYTES,
meta.transfer_stats.store.bytes,
)
transfer_stats.increase_counter(
_TransferMetricName.STORE_TIME,
meta.transfer_stats.store.time,
)
for size in meta.transfer_stats.store.sizes:
transfer_stats.observe_histogram(
_TransferMetricName.STORE_SIZE, size
)
self._connector_stats.aggregate(transfer_stats)
for job_id, count in meta.completed_jobs.items():
assert count > 0
if job_id < self._stale_job_threshold:
logger.debug(
"Skipping stale completed job %d (pre-reset counter: %d)",
job_id,
self._stale_job_threshold,
)
continue
job_status = self._jobs[job_id]
job_status.pending_count -= count
if job_status.pending_count > 0:
continue
assert job_status.pending_count == 0
req_status = self._req_status[job_status.req_id]
if job_status.is_store:
self.manager.complete_store(job_status.keys, req_status.req_context)
else:
self.manager.complete_load(job_status.keys, req_status.req_context)
if self._chunks_being_loaded:
self._chunks_being_loaded.difference_update(job_status.keys)
if self._block_id_to_pending_jobs:
# Sliding window blocks are tracked from store creation
# and must be cleaned up unconditionally.
self._remove_pending_job(job_id, job_status.fenced_block_ids)
# Non-sliding-window blocks are only tracked after
# request_finished, so only clean up for finished requests.
if req_status.req.is_finished():
self._remove_pending_job(
job_id, job_status.deferred_fence_block_ids
)
del self._jobs[job_id]
req_status.transfer_jobs.remove(job_id)
if req_status.finished_signaled and not req_status.transfer_jobs:
del self._req_status[job_status.req_id]
def get_stats(self) -> OffloadingConnectorStats | None:
stats: OffloadingConnectorStats | None = None
if not self._connector_stats.is_empty():
stats = self._connector_stats
self._connector_stats = OffloadingConnectorStats()
manager_stats = self.manager.get_stats()
if manager_stats is not None:
if stats is None:
stats = manager_stats
else:
stats.aggregate(manager_stats)
return stats
def request_finished(
self,
request: Request,
) -> tuple[bool, dict[str, Any] | None]:
"""Called when a request has finished, before its blocks are freed.
Returns:
True if the request is being saved/sent asynchronously and blocks
should not be freed until the request_id is returned from
get_finished().
Optional KVTransferParams to be included in the request outputs
returned by the engine.
"""
req_status = self._req_status.get(request.request_id)
if req_status is None:
# Untracked request (offloading never started): no in-flight jobs,
# nothing was deferred, so finalize immediately.
req_context = _create_req_context(request)
self.manager.on_new_request(req_context)
self.manager.on_request_finished(req_context)
return False, None
self._maybe_observe_lookup_async_delay(req_status)
# Update offload keys with final block hash so _build_store_jobs can
# create store jobs for the last block(s) on the next schedule step.
req_status.update_offload_keys()
# Keep req_status alive: _build_store_jobs will process finished_req_ids
# on the next step and handle cleanup after creating store jobs.
# Register deferred fences so future block reuse triggers a flush via
# _block_id_to_pending_jobs.
for job_id in req_status.transfer_jobs:
job_status = self._jobs[job_id]
for bid in job_status.deferred_fence_block_ids or ():
self._block_id_to_pending_jobs.setdefault(bid, set()).add(job_id)
return False, None
def take_events(self) -> Iterable[KVCacheEvent]:
"""Drain pending KV cache events.
Complete metadata is available only when self-describing KV events
are enabled, and only for full-attention groups. Other shapes retain
the previous placeholder payload so consumers can ignore them.
Yields:
``BlockStored`` or ``BlockRemoved`` events corresponding to
the underlying :class:`OffloadingEvent` stream.
"""
yield from self._events_tracker.take_events(self.manager.take_events())
def reset_cache(self) -> None:
"""Reset the offloading manager cache, evicting all stored chunks."""
# reset_cache cannot be called in the middle of a schedule step
assert not self._current_batch_load_jobs
assert not self._current_batch_jobs_to_flush
assert not self._current_batch_allocated_block_ids
# Flush all in-flight jobs
self._current_batch_jobs_to_flush.update(self._jobs.keys())
for req_id, status in list(self._req_status.items()):
if status.req.is_finished():
if not status.finished_signaled:
self.manager.on_request_finished(status.req_context)
del self._req_status[req_id]
# Reset offloading manager cache
self.manager.reset_cache()
# Reset store progress so active requests re-offload from chunk 0.
for status in self._req_status.values():
for group_state in status.group_states:
group_state.next_stored_chunk_idx = 0
status.transfer_jobs.clear()
status.partial_tail_boundary = None
# Discard jobs and save job_counter to be able to discard worker responses
self._stale_job_threshold = self._job_counter
self._jobs.clear()
self._block_id_to_pending_jobs.clear()
# The manager pool is empty; pending event payloads and announced
# reference counts are stale.
self._events_tracker.reset()
# Note: _current_batch_jobs_to_flush is intentionally NOT cleared.
# The load flush IDs collected above must be delivered to workers.
if self._chunks_being_loaded is not None:
self._chunks_being_loaded.clear()
def shutdown(self) -> None:
self.manager.shutdown()