Skip to content

vllm_omni.core.prefix_cache

Omni prefix cache.

Structure mirrors vllm/v1/core.

Import invariant: this package must not import vllm at module scope (TYPE_CHECKING and function-body imports are fine). tests/core/test_prefix_cache.py loads these modules without a vllm install; a top-level from vllm... turns that suite red.

That is also why these files use logging.getLogger instead of the repo-wide init_logger — init_logger comes from vllm.logger. The deviation is required, not an oversight.

_merge_uncached_mm is the exception at call time: it lazily imports vllm_omni.utils.mm_outputs, which pulls in vllm. Import of this package still succeeds without vllm; materialize of leftover mm does not.

Modules:

Name Description
block_pool

Pinned CPU block mirror (mirrors vLLM BlockPool).

controller

Runs WriteTasks and the step device→host staging pool.

group_view

KV-cache group access for the omni prefix cache.

interface

Interface types for the omni prefix cache.

manager

Omni prefix cache, manager side.

runner_mixin

Runner-facing access layer for the omni prefix cache.

HIDDEN_KEY module-attribute

HIDDEN_KEY: TensorName = '__hidden_states__'

FullAttentionGroupView

View over the first (full-attention) KV-cache group.

Step slots come from the CPU block table (step_slots_cpu), not the device slot_mapping.

block_size instance-attribute

block_size = block_size

batch_req_ids

batch_req_ids() -> list[str]

step_slots_cpu

step_slots_cpu(
    req_ids: list[str], num_scheduled: dict[str, int]
) -> Tensor

This step's slot mapping, computed on CPU from the block table.

The device slot_mapping would need a stream sync to read back, which stalls the whole forward; the CPU block table carries the same information (positions are num_computed .. +num_scheduled per request).

ModelCachePolicy dataclass

Replaces getattr probing on models for cache behavior decisions.

Hidden's name is HIDDEN_KEY (shared identity). This object answers whether this model caches it, and which mm keys are deferred.

deferred_keys class-attribute instance-attribute

deferred_keys: frozenset[TensorName] = frozenset()

hidden_key property

hidden_key: TensorName | None

Pool key for hidden, or None when this model opts out.

needs_full_hidden_states class-attribute instance-attribute

needs_full_hidden_states: bool = True

from_model classmethod

from_model(model: Any) -> ModelCachePolicy

Shim over legacy per-model attributes (deprecation window).

get_hit_keys

get_hit_keys(
    keys: Iterable[TensorName],
) -> list[TensorName]

Keys to plan/prefetch for a hit: hidden first (if cached), then mm.

skip_immediate_mm

skip_immediate_mm(key: TensorName) -> bool

Immediate on-device clone must not take hidden or deferred keys from mm.

OmniPrefixCacheController

Staging pool + committer. Step device→host is launched at save; this thread waits that event (JOIN_NEXT_STEP) or copies deferred rows (JOIN_ON_FINISH), then writes into the CPU pool.

append_chunk

append_chunk(
    task: WriteTask,
    chunk: _WriteChunk,
    freeze_event: Event | None = None,
) -> TaskState | None

Append to a pending task. None when appended, else the closing state.

dispatch

dispatch(tasks: list[WriteTask]) -> None

Hand registered, queued tasks to the copy path. Threaded: enqueue. Eager: the copy + pool write run here, inline — never call this under the manager's state lock.

drain_completed

drain_completed() -> list[int]

Pop pool-written tasks from _completed and drop them from _tasks. Staging holders were already released at the pool write.

drain_failed

drain_failed() -> list[int]

Pop failed task ids from _failed. Does not drop _tasks.

escalate

escalate(tids: list[int]) -> None

Move pending tasks to the front of the high-priority queue.

PENDING -> QUEUED at the head; QUEUED on the low-priority queue moves up; QUEUED already high-priority is a no-op. COPYING and later belong to the worker and are untouched: the worker claims under _wake too, so a task cannot be popped and re-queued behind its back.

fetch_host

fetch_host(
    task: WriteTask, slots: Tensor, key: str
) -> Tensor

Rows for slots of one not-yet-done JOIN_ON_FINISH task.

_slot_ref puts JOIN_NEXT_STEP tids in join_tids (wait then pool). This path reads committer-written chunk.host, or the device clone if that device→host has not landed.

get_task

get_task(tid: int) -> WriteTask | None

in_flight_tasks

in_flight_tasks() -> int

Registered tasks not yet drained (diagnostics only).

join

join(tids: list[int]) -> None

Block until each task has finished the CPU-pool write (or failed).

join_host_ready

join_host_ready(tids: list[int]) -> None

Block until each task's device→host is complete (host_ready).

Staging: committer has waited step_d2h_event. Deferred: committer has written chunk.host. Does not wait for the CPU-pool write.

pin_budget

pin_budget(ticket: _BudgetTicket, tid: int) -> None

Record that tid holds a view of the clone ticket charges.

register

register(task: WriteTask, queued: bool = True) -> None

Make a task visible (registry + QUEUED) without running anything; safe under the manager's state lock. queued=False (deferred tasks) stays PENDING on the GPU clone until finish/abort or the GPU-byte budget forces a copy. Queued tasks must then go through dispatch.

Caller must reserve() the task bytes first (budget flush can block; the manager does that outside the state lock) and pin the task on its budget ticket(s) before register.

reserve

reserve(nbytes: int) -> None

Reserve GPU-clone bytes; blocking flush happens here, so callers must not hold the manager's state lock.

shutdown

shutdown() -> None

stage_step_host

stage_step_host(
    tensors: dict[str, Tensor],
    n: int,
    freeze_event: Event | None,
    step_holder: StagingBufferHolder,
) -> StepD2HClaim

Claim a staging slot and, when tensors is non-empty, launch ONE whole-step device→host into it.

Leftover-only saves pass empty tensors and still take a slot (empty views) so every step id shares this bound. A full pool waits for materialize/discard; timeout then errors. A step larger than the page overflows the next slot — that is a config break. The caller binds tasks after submit; step_holder is released by materialize/discard via staging_release.

staging_bind

staging_bind(
    slot: int, holder: StagingBufferHolder
) -> None

staging_release

staging_release(
    slot: int, holder: StagingBufferHolder
) -> None

submit

submit(task: WriteTask, queued: bool = True) -> None

register + dispatch in one call (callers not holding the state lock).

OmniPrefixCacheManager

discard_step

discard_step(step_id: int) -> None

Consume the step context when nothing will materialize it.

Any thread; same exactly-once contract as materialize (unknown or duplicate id raises). Only the read-side snapshot is dropped — the cache write proceeds unchanged.

materialize

materialize(
    step_id: int, req_ids: list[str]
) -> StageCacheOutputs

Per-request merged outputs for the step saved as step_id.

Any thread. req_ids must be (a subset of) the save-time snapshot; an outside id means the caller is reading the live batch. A request without a hit is a plain miss and gets exactly this step's rows — normal path, nothing logged. A hit span that resolves to absent rows raises OmniPrefixCacheUnmatchError: fatal by contract (do not pretend it was a miss).

Two phases: under the lock, publish finished writes and pin every row source (task refs + masks, absent checks included) — not yet reading the tensors. Unlocked: wait this step's step_d2h_event, clone the staging views (then drop the step holder), and merge. The engine thread never waits on this thread's device→host copy.

new_step_starts

new_step_starts(scheduler_output: SchedulerOutput) -> None

Handle one scheduler_output.

Engine thread only; before _update_states removes finished requests; exactly once per real step. Registers new-request prefix hits (copying their block tables) and forces finished/aborted requests' still-open writes onto the high-priority copy queue — a block hash that entered the batch must land in the cache, abort included. escalate (eager: the copy + pool write) runs after _state_lock is released.

register_policy

register_policy(policy: ModelCachePolicy) -> None

save_outputs

save_outputs(
    hidden_states: Tensor | None,
    mm_outputs: dict[str, Any] | None,
    *,
    num_tokens_unpadded: int,
    num_tokens_padded: int,
) -> int

Write this step's outputs into the cache; returns the step id.

Engine thread only, after the forward and before materialize. Immediately-cached rows: one on-device clone, one whole-step device→host into the staging pool, then one JOIN_NEXT_STEP WriteTask per request whose chunk.host is a view of that page. Deferred rows stay on the device clone (JOIN_ON_FINISH); the committer copies them later. Leftover mm (this-step deferred rows + mm not written to the pool) is copied to CPU here so materialize never reads live graph buffers. Snapshots everything materialize needs. The returned step id MUST be consumed exactly once — by materialize() or discard_step(). Every step id claims one staging slot (saves with only leftover mm included); a later save waits for a free slot and times out if none return.

The state lock never covers a blocking wait: the previous step's JOIN_NEXT_STEP wait, the clone build, the GPU-byte-budget reserve (which may flush), and the staging-slot claim all run unlocked.

shutdown

shutdown() -> None

OmniPrefixCacheStagingTimeoutError

Bases: OmniPrefixCacheUnmatchError

Save waited for a free staging slot and timed out.

OmniPrefixCacheUnmatchError

Bases: RuntimeError

Contract, config, or KV-occupancy error that must raise.

Includes hit spans that resolve to absent slots (omni cache diverged from vLLM KV), a step id consumed twice or never saved, a step larger than the staging page, and a failed write. Do not pretend these were a miss.

PrefixBlockPool

Durable CPU mirror of vLLM KV slots, one tensor per name.

Storage is (num_blocks, block_size, feat) per key, viewed flat as (num_slots, feat) so vLLM slot ids index rows directly. Write access is single-writer (the controller committer thread); readers take row views.

_caches is not an unbounded store and has no LRU. Dict keys are tensor names (__hidden_states__ plus mm whose first dim is tokens), opened once by ensure_key and never dropped — a model emits a handful of those names, not a per-request set. Each value is a fixed tensor sized to the same num_blocks as the upstream KV pool. Rows are those kv slots: when vLLM recycles a block, the next pool write overwrites that row. Slot reuse is the eviction.

alloc_key

alloc_key(
    key: str, dtype: dtype, feat: int
) -> Tensor | None

Allocate storage for a new key without publishing it (None if already open). Pinned allocation is slow; run this unlocked and hand the tensor to install_key under the manager's state lock.

ensure_key

ensure_key(key: str, dtype: dtype, feat: int) -> None

has_key

has_key(key: str) -> bool

install_key

install_key(key: str, storage: Tensor) -> None

keys

keys() -> set[TensorName]

rows

rows(key: str, slots: Tensor) -> Tensor

Gather rows for non-contiguous slots (returns a copy).

write

write(key: str, slots: Tensor, src_cpu: Tensor) -> None

Write rows into the pool; caller (committer thread) is the single writer.

PrefixCacheConfig dataclass

Sizing and flow-control knobs (mirrors KVCacheConfig).

block_size instance-attribute

block_size: int

copy_chunk_bytes class-attribute instance-attribute

copy_chunk_bytes: int = 16 * 1024 * 1024

gpu_staging_bytes class-attribute instance-attribute

gpu_staging_bytes: int = 512 * 1024 * 1024

num_blocks instance-attribute

num_blocks: int

staging_capacity_tokens class-attribute instance-attribute

staging_capacity_tokens: int = 1024

staging_claim_timeout_s class-attribute instance-attribute

staging_claim_timeout_s: float = 30.0

staging_depth class-attribute instance-attribute

staging_depth: int = 4

from_vllm_config classmethod

from_vllm_config(
    *,
    num_blocks: int,
    block_size: int,
    scheduler_config: Any = None,
    model_config: Any = None,
) -> PrefixCacheConfig

Size device→host staging from the running scheduler.

A slot holds one step (the whole batch), not one request: staging_capacity_tokens is max_num_batched_tokens (falls back to max_model_len); staging_depth is how many unconsumed step ids may exist at once, not max_num_seqs.

Pinned staging is allocated lazily per key at depth * capacity_tokens * width * dtype. There is no clamp: a step larger than capacity raises. A 16k-token thinking batch at hidden=2048 bf16 is ~256 MiB for hidden alone; each mm key adds another pool tensor.

staging_depth is the dataclass default (4). There is no CLI or deploy YAML knob — changing it is a code change. Every save that issues a step id claims one slot, including saves with only leftover mm. A full pool waits for materialize/discard; staging_claim_timeout_s then errors.

StageCacheOutputs

Bases: NamedTuple

Plain value object: per-request merged stage outputs.

hidden_states instance-attribute

hidden_states: dict[ReqId, Tensor] | None

mm_outputs instance-attribute

mm_outputs: dict[TensorName, dict[ReqId, Any]]

WriteSchedule

Bases: Enum

Write scheduling policy for one WriteTask.

JOIN_NEXT_STEP class-attribute instance-attribute

JOIN_NEXT_STEP = 'join_next_step'

JOIN_ON_FINISH class-attribute instance-attribute

JOIN_ON_FINISH = 'join_on_finish'

check_prefix_cache_kv_groups

check_prefix_cache_kv_groups(
    kv_cache_groups: object,
) -> None

Reject hybrid / multi-group models at kv-cache init, not first step.

Only needs kv_cache_config.kv_cache_groups. FullAttentionSpec is imported here so tests/core can load this module without vllm.

check_prefix_cache_kv_transfer

check_prefix_cache_kv_transfer(
    kv_transfer_config: object,
) -> None

Reject kv_consumer / kv_both stages.

KV loaded from a producer also shows up as num_computed_tokens; the manager would read it as a local hit with no rows behind it.

get_prefix_cache_group_view

get_prefix_cache_group_view(
    input_batch: InputBatch,
    block_size: int,
    kv_cache_groups: object = None,
) -> FullAttentionGroupView | None

Build the group view; None only if the batch has no block table.

Group spec is checked first (raises). Selection is by spec, not by counting block tables: a hybrid model's group 0 is not necessarily full attention, and a narrower per-group table would make step_slots_cpu silently clamp.

stage_prefix_cache_config

stage_prefix_cache_config(
    *,
    kv_cache_config: object,
    cache_config: object,
    kv_transfer_config: object,
    scheduler_config: object,
    model_config: object,
    is_pooling_model: bool,
    speculative_config: object = None,
) -> PrefixCacheConfig | None

Runner-side gate shared by the GPU and NPU model runners.

Returns None when the stage does not run an omni prefix cache (enable_prefix_caching off, or a pooling stage that never saves). Otherwise refuses kv_consumer / kv_both, speculative decoding, sub-block matching and hybrid kv groups, then sizes the config from the scheduler. One place so a platform runner cannot silently skip a refusal the other one has.