def build_offloading_config(
vllm_config: "VllmConfig",
kv_cache_config: "KVCacheConfig",
) -> OffloadingConfig:
"""Translate vLLM configuration into the native offloading boundary."""
kv_transfer_config = vllm_config.kv_transfer_config
assert kv_transfer_config is not None
extra_config = kv_transfer_config.kv_connector_extra_config
assert kv_transfer_config.engine_id is not None
engine_id = kv_transfer_config.engine_id
parallel_config = vllm_config.parallel_config
selected_groups = tuple(
(group_id, kv_cache_config.kv_cache_groups[group_id])
for group_id in get_offloading_group_ids(kv_cache_config)
)
if not selected_groups:
raise ValueError("KV offloading found no eligible cache groups.")
groups = tuple(
OffloadingGroupConfig(
group_id=group_id,
tokens_per_block=resolve_dcp_kv_block_size(
group.kv_cache_spec,
parallel_config.decode_context_parallel_size,
),
layer_names=tuple(group.layer_names),
)
for group_id, group in selected_groups
)
_, tokens_per_hash = resolve_kv_cache_block_sizes(kv_cache_config, vllm_config)
for group in groups:
assert group.tokens_per_block % tokens_per_hash == 0, (
f"tokens_per_block={group.tokens_per_block} not divisible by "
f"tokens_per_hash={tokens_per_hash}. "
f"Hybrid models (e.g. Mamba+Attention) need "
f"--enable-prefix-caching to align block sizes."
)
blocks_per_chunk = 1
blocks_per_chunk_config = extra_config.get("blocks_per_chunk")
tokens_per_chunk = extra_config.get("block_size")
if blocks_per_chunk_config is not None and tokens_per_chunk is not None:
raise ValueError(
"Specify only one of 'block_size' or 'blocks_per_chunk' "
"in kv_connector_extra_config."
)
if blocks_per_chunk_config is not None:
blocks_per_chunk = int(blocks_per_chunk_config)
if blocks_per_chunk <= 0:
raise ValueError("'blocks_per_chunk' must be greater than 0.")
elif tokens_per_chunk is not None:
tokens_per_chunk_int = int(tokens_per_chunk)
unique_tokens_per_block = {group.tokens_per_block for group in groups}
assert len(unique_tokens_per_block) == 1, (
"If 'block_size' is specified in kv_connector_extra_config, "
"there must be at least one KV cache group, "
"and all groups must have the same block size."
)
tokens_per_block = unique_tokens_per_block.pop()
if tokens_per_chunk_int % tokens_per_block == 0:
blocks_per_chunk = tokens_per_chunk_int // tokens_per_block
else:
raise ValueError(
f"'block_size'={tokens_per_chunk_int} in kv_connector_extra_config "
f"must be a multiple of the GPU KV cache block size "
f"({tokens_per_block} tokens). Use "
f"{round_up(tokens_per_chunk_int, tokens_per_block)} instead, or set "
f"'blocks_per_chunk' to express the chunk size in blocks."
)
worker_kv_bytes_per_block = 0
if (
kv_cache_config.hisparse_host_num_blocks is None
and kv_cache_config.num_blocks > 0
and kv_cache_config.kv_cache_tensors
):
# Scratch filtering must preserve the scheduler/worker allocation stride.
# Every KVCacheTensor describes placement within the same backing allocation,
# so its size is the total, not a per-tensor share.
total_gpu_kv_bytes = kv_cache_config.kv_cache_tensors[0].size
worker_kv_bytes_per_block = total_gpu_kv_bytes // kv_cache_config.num_blocks
elif kv_cache_config.num_blocks > 0:
worker_kv_bytes_per_block = sum(
_group_kv_bytes_per_block(group) for _, group in selected_groups
)
single_group_spec = (
kv_cache_config.kv_cache_groups[0].kv_cache_spec
if len(kv_cache_config.kv_cache_groups) == 1
else None
)
replicated_layout = (
vllm_config.model_config.use_mla
and parallel_config.tensor_parallel_size > 1
and kv_cache_config.kv_tp_replicas == parallel_config.tensor_parallel_size
and worker_kv_bytes_per_block > 0
# Safe MVP boundary: TP-only, no other parallel axes.
and parallel_config.pipeline_parallel_size == 1
and parallel_config.prefill_context_parallel_size == 1
and parallel_config.decode_context_parallel_size == 1
and parallel_config.world_size == parallel_config.tensor_parallel_size
# Shared /dev/shm mmap layout is single-node mp only.
and parallel_config.distributed_executor_backend == "mp"
and parallel_config.nnodes_within_dp == 1
)
canonical_layout = bool(extra_config.get("canonical_layout", False))
# Only a single non-MLA full-attention group with genuinely head-sharded
# pages is parallelism-invariant: replicated latent or GQA heads,
# per-token-head scales, CP token sharding, and the V2 model runner's
# layout are all excluded.
is_parallelism_agnostic = (
not vllm_config.use_v2_model_runner
and single_group_spec is not None
and isinstance(single_group_spec, FullAttentionSpec)
and not isinstance(single_group_spec, MLAAttentionSpec)
and single_group_spec.num_kv_heads * parallel_config.tensor_parallel_size
== vllm_config.model_config.get_total_num_kv_heads()
and not single_group_spec.kv_quant_mode.is_per_token_head
and parallel_config.decode_context_parallel_size == 1
and parallel_config.prefill_context_parallel_size == 1
)
# Canonical pages are topology-free, so the gate widens to every config
# whose mappings derive portable, group by group; certification happens
# per layer at registration and create_worker fails closed on this flag.
if canonical_layout and not is_parallelism_agnostic:
tp_size = parallel_config.tensor_parallel_size
total_kv_heads = vllm_config.model_config.get_total_num_kv_heads()
def spec_certifiable(spec: KVCacheSpec) -> bool:
"""Conservative static mirror of _layer_mapping's per-layer checks."""
if not isinstance(spec, AttentionSpec):
return False
if spec.kv_quant_mode.is_per_token_head:
return False
if type(spec) is MLAAttentionSpec:
return (
spec.tokens_per_state == 1
and spec.real_page_size_bytes % spec.block_size == 0
)
if isinstance(spec, (SlidingWindowMLASpec, MLAAttentionSpec)):
return False
if not isinstance(spec, (FullAttentionSpec, SlidingWindowSpec)):
return False
return (
total_kv_heads % tp_size == 0 or tp_size % total_kv_heads == 0
) and spec.num_kv_heads == max(1, total_kv_heads // tp_size)
# UniformTypeKVCacheSpecs groups (e.g. MLA plus its DSA indexer) hold
# one spec per layer; certify per layer, as the mapping derivation does.
layer_specs = [
spec
for group in kv_cache_config.kv_cache_groups
for spec in iter_layer_specs(group.kv_cache_spec)
]
is_parallelism_agnostic = (
len(layer_specs) > 0
and all(spec_certifiable(spec) for spec in layer_specs)
and parallel_config.decode_context_parallel_size == 1
and parallel_config.prefill_context_parallel_size == 1
and parallel_config.world_size == tp_size
)
if canonical_layout:
replicated_layout = (
is_parallelism_agnostic
and all(
type(spec) is MLAAttentionSpec
for _, group in selected_groups
for spec in iter_layer_specs(group.kv_cache_spec)
)
and parallel_config.nnodes_within_dp == 1
and (
parallel_config.world_size == 1
or parallel_config.distributed_executor_backend == "mp"
)
)
kv_events_config = vllm_config.kv_events_config
cache_dtype = (
vllm_config.model_config.dtype
if vllm_config.cache_config.cache_dtype == "auto"
else vllm_config.cache_config.cache_dtype
)
return OffloadingConfig(
groups=groups,
worker_kv_bytes_per_block=worker_kv_bytes_per_block,
enable_kv_cache_events=(
kv_events_config is not None and kv_events_config.enable_kv_cache_events
),
extra_config=extra_config,
engine_id=engine_id,
model=OffloadingModelConfig(
name=vllm_config.model_config.model,
dtype=str(cache_dtype).removeprefix("torch."),
),
cache=OffloadingCacheConfig(
tokens_per_hash=tokens_per_hash,
blocks_per_chunk=blocks_per_chunk,
),
parallel=OffloadingParallelConfig(
rank=parallel_config.rank,
world_size=parallel_config.world_size,
tp_size=parallel_config.tensor_parallel_size,
pp_size=parallel_config.pipeline_parallel_size,
pcp_size=parallel_config.prefill_context_parallel_size,
dcp_size=parallel_config.decode_context_parallel_size,
data_parallel_index=parallel_config.data_parallel_index,
data_parallel_size=parallel_config.data_parallel_size,
data_parallel_rank_local=parallel_config.data_parallel_rank_local,
is_parallelism_agnostic=is_parallelism_agnostic,
),
replicated_layout=replicated_layout,
canonical_layout=canonical_layout,
kv_cache_layout=vllm_config.cache_config.kv_cache_layout,
)