[DCP] Match the replicated draft KV pool's page granularity to its allocator (#33348)

This commit is contained in:
Khoa Pham
2026-08-05 11:40:28 -07:00
committed by GitHub
parent b1bd871df5
commit 5f79cf3511
@@ -270,6 +270,19 @@ class KVCacheConfigurator:
unified_memory_pool=pools.unified_memory_pool,
)
# Note(kpham-sgl):
# 1. A replicated draft indexes the allocator's virtual locs raw, so its pools
# span and page that space; the sharded target translates and stays per-rank.
# 2. A pool must page as its allocator does, or its last page falls short.
@property
def loc_space_scale(self) -> int:
dcp_size = self.server_args.dcp_size
return dcp_size if (self.is_draft_worker and dcp_size > 1) else 1
@property
def pool_page_size(self) -> int:
return get_schedule().page_size * self.loc_space_scale
def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes:
max_total_num_tokens = config.max_total_num_tokens
max_running_requests = config.max_running_requests
@@ -281,13 +294,12 @@ class KVCacheConfigurator:
# Draft pools are replicated, not DCP-sharded, yet consume the shared
# allocator's virtual locs in [0, max_total * dcp_size) untranslated.
dcp_size = self.server_args.dcp_size
if self.is_draft_worker and dcp_size > 1:
max_total_num_tokens *= dcp_size
loc_scale = self.loc_space_scale
max_total_num_tokens *= loc_scale
if full_max_total_num_tokens is not None:
full_max_total_num_tokens *= dcp_size
full_max_total_num_tokens *= loc_scale
if swa_max_total_num_tokens is not None:
swa_max_total_num_tokens *= dcp_size
swa_max_total_num_tokens *= loc_scale
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
# full/swa sizes but does NOT own c4/c128/state pools (those live on
@@ -1000,7 +1012,7 @@ class KVCacheConfigurator:
c128_size=c128_max_total_num_tokens,
c4_state_pool_size=c4_state_pool_size,
c128_state_pool_size=c128_state_pool_size,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
swa_page_size=swa_page_size,
sliding_window=self.model_config.window_size,
dtype=self.kv_cache_dtype,
@@ -1026,7 +1038,7 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_dsa_kv_pool_cls()
token_to_kv_pool = PoolCls(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1050,7 +1062,7 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_mla_kv_pool_cls()
token_to_kv_pool = PoolCls(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1067,7 +1079,7 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_mha_kv_pool_cls()
token_to_kv_pool = PoolCls(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
@@ -1104,7 +1116,7 @@ class KVCacheConfigurator:
token_to_kv_pool = SWAKVPool(
size=full_max_total_num_tokens,
size_swa=swa_max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
post_capture_active=self.post_capture_kv_active,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1126,7 +1138,7 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMLATokenToKVPool(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1146,7 +1158,7 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMHATokenToKVPool(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
@@ -1186,7 +1198,7 @@ class KVCacheConfigurator:
PoolCls = DSATokenToKVPool
token_to_kv_pool = PoolCls(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1208,7 +1220,7 @@ class KVCacheConfigurator:
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MLATokenToKVPoolFP4(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1223,7 +1235,7 @@ class KVCacheConfigurator:
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MLATokenToKVPool(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1286,7 +1298,7 @@ class KVCacheConfigurator:
token_to_kv_pool = SWAKVPool(
size=full_max_total_num_tokens,
size_swa=size_swa,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
post_capture_active=self.post_capture_kv_active,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1311,7 +1323,7 @@ class KVCacheConfigurator:
)
token_to_kv_pool = MiniMaxSparseKVPool(
size=max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
# fp8 attn-GEMM mode opts the lightning-indexer cache into
# fp8 too (fp8 indexer GEMMs); fp8 KV without the mode
@@ -1368,7 +1380,7 @@ class KVCacheConfigurator:
else mha_pool_class
)
token_to_kv_pool = HybridLinearKVPool(
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
size=max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1391,7 +1403,7 @@ class KVCacheConfigurator:
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MHATokenToKVPoolFP4(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
@@ -1424,7 +1436,7 @@ class KVCacheConfigurator:
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
token_to_kv_pool = pool_cls(
max_total_num_tokens,
page_size=get_schedule().page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,