[DCP] Match the replicated draft KV pool's page granularity to its allocator (#33348)
This commit is contained in:
@@ -270,6 +270,19 @@ class KVCacheConfigurator:
|
|||||||
unified_memory_pool=pools.unified_memory_pool,
|
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:
|
def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes:
|
||||||
max_total_num_tokens = config.max_total_num_tokens
|
max_total_num_tokens = config.max_total_num_tokens
|
||||||
max_running_requests = config.max_running_requests
|
max_running_requests = config.max_running_requests
|
||||||
@@ -281,13 +294,12 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
# Draft pools are replicated, not DCP-sharded, yet consume the shared
|
# Draft pools are replicated, not DCP-sharded, yet consume the shared
|
||||||
# allocator's virtual locs in [0, max_total * dcp_size) untranslated.
|
# allocator's virtual locs in [0, max_total * dcp_size) untranslated.
|
||||||
dcp_size = self.server_args.dcp_size
|
loc_scale = self.loc_space_scale
|
||||||
if self.is_draft_worker and dcp_size > 1:
|
max_total_num_tokens *= loc_scale
|
||||||
max_total_num_tokens *= dcp_size
|
if full_max_total_num_tokens is not None:
|
||||||
if full_max_total_num_tokens is not None:
|
full_max_total_num_tokens *= loc_scale
|
||||||
full_max_total_num_tokens *= dcp_size
|
if swa_max_total_num_tokens is not None:
|
||||||
if swa_max_total_num_tokens is not None:
|
swa_max_total_num_tokens *= loc_scale
|
||||||
swa_max_total_num_tokens *= dcp_size
|
|
||||||
|
|
||||||
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
|
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
|
||||||
# full/swa sizes but does NOT own c4/c128/state pools (those live on
|
# 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,
|
c128_size=c128_max_total_num_tokens,
|
||||||
c4_state_pool_size=c4_state_pool_size,
|
c4_state_pool_size=c4_state_pool_size,
|
||||||
c128_state_pool_size=c128_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,
|
swa_page_size=swa_page_size,
|
||||||
sliding_window=self.model_config.window_size,
|
sliding_window=self.model_config.window_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
@@ -1026,7 +1038,7 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = current_platform.get_dsa_kv_pool_cls()
|
PoolCls = current_platform.get_dsa_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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()
|
PoolCls = current_platform.get_mla_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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()
|
PoolCls = current_platform.get_mha_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
@@ -1104,7 +1116,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool = SWAKVPool(
|
token_to_kv_pool = SWAKVPool(
|
||||||
size=full_max_total_num_tokens,
|
size=full_max_total_num_tokens,
|
||||||
size_swa=swa_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,
|
dtype=self.kv_cache_dtype,
|
||||||
post_capture_active=self.post_capture_kv_active,
|
post_capture_active=self.post_capture_kv_active,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
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(
|
token_to_kv_pool = NPUMLATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
@@ -1146,7 +1158,7 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
token_to_kv_pool = NPUMHATokenToKVPool(
|
token_to_kv_pool = NPUMHATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
@@ -1186,7 +1198,7 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = DSATokenToKVPool
|
PoolCls = DSATokenToKVPool
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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:
|
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MLATokenToKVPoolFP4(
|
token_to_kv_pool = MLATokenToKVPoolFP4(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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:
|
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MLATokenToKVPool(
|
token_to_kv_pool = MLATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
@@ -1286,7 +1298,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool = SWAKVPool(
|
token_to_kv_pool = SWAKVPool(
|
||||||
size=full_max_total_num_tokens,
|
size=full_max_total_num_tokens,
|
||||||
size_swa=size_swa,
|
size_swa=size_swa,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
post_capture_active=self.post_capture_kv_active,
|
post_capture_active=self.post_capture_kv_active,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
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(
|
token_to_kv_pool = MiniMaxSparseKVPool(
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
# fp8 attn-GEMM mode opts the lightning-indexer cache into
|
# fp8 attn-GEMM mode opts the lightning-indexer cache into
|
||||||
# fp8 too (fp8 indexer GEMMs); fp8 KV without the mode
|
# fp8 too (fp8 indexer GEMMs); fp8 KV without the mode
|
||||||
@@ -1368,7 +1380,7 @@ class KVCacheConfigurator:
|
|||||||
else mha_pool_class
|
else mha_pool_class
|
||||||
)
|
)
|
||||||
token_to_kv_pool = HybridLinearKVPool(
|
token_to_kv_pool = HybridLinearKVPool(
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
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:
|
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MHATokenToKVPoolFP4(
|
token_to_kv_pool = MHATokenToKVPoolFP4(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
@@ -1424,7 +1436,7 @@ class KVCacheConfigurator:
|
|||||||
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
|
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
|
||||||
token_to_kv_pool = pool_cls(
|
token_to_kv_pool = pool_cls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=get_schedule().page_size,
|
page_size=self.pool_page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
|
|||||||
Reference in New Issue
Block a user