[AMD] Enable unified-KV HiCache on DeepSeek-V4 (#29417)
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -1163,6 +1163,16 @@ class Req(ReqDllmMixin):
|
||||
token_ids_to_match = self.full_untruncated_fill_ids
|
||||
key_limit: Optional[int] = self._compute_max_prefix_len(input_len)
|
||||
|
||||
# SWA lives in a per-request ring that's not content-stable and is never
|
||||
# stored in the radix tree, so a reused prefix carries stale SWA. Cap the
|
||||
# match by the trailing sliding window so it gets re-prefilled, rewriting
|
||||
# this request's SWA ring. No-op for other layouts.
|
||||
if tree_cache is not None:
|
||||
reprefill_tail = tree_cache.swa_reprefill_tail_tokens()
|
||||
if reprefill_tail:
|
||||
capped = max(0, input_len - reprefill_tail)
|
||||
key_limit = capped if key_limit is None else min(key_limit, capped)
|
||||
|
||||
# Disable prefix caching when embed overrides are present: same token IDs
|
||||
# with different override vectors must not share cached KV values.
|
||||
if self.positional_embed_overrides is not None:
|
||||
|
||||
@@ -100,9 +100,10 @@ def match_prefix_for_req(
|
||||
if token_ids is None:
|
||||
token_ids = req.origin_input_ids + req.output_ids
|
||||
|
||||
# unified_kv SWA lives in a per-request ring (not content-stable, never cached
|
||||
# in the radix tree), so a reused prefix carries stale SWA. Cap the match by the
|
||||
# trailing sliding window so it is re-prefilled. No-op for other layouts.
|
||||
# unified_kv SWA lives in a per-request ring that's not content-stable and is
|
||||
# never stored in the radix tree, so a reused prefix carries stale SWA. Cap
|
||||
# the match by the trailing sliding window so it gets re-prefilled, rewriting
|
||||
# this request's SWA ring. No-op for other layouts.
|
||||
reprefill_tail = tree_cache.swa_reprefill_tail_tokens()
|
||||
key_limit = max(0, len(token_ids) - reprefill_tail) if reprefill_tail else None
|
||||
|
||||
|
||||
@@ -329,6 +329,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
|
||||
return False
|
||||
|
||||
def swa_reprefill_tail_tokens(self) -> int:
|
||||
# Only the unified_kv compress-only HiCache layout needs to hold back a
|
||||
# trailing sliding window for re-prefill; every other cache keeps SWA
|
||||
# content-stable and overrides this where relevant.
|
||||
return 0
|
||||
|
||||
def supports_mamba(self) -> bool:
|
||||
|
||||
@@ -390,7 +390,7 @@ class DeepSeekV4LayerItem(NamedTuple):
|
||||
class DeepSeekV4UnifiedKVPool:
|
||||
"""
|
||||
Layout:
|
||||
unified_kv[L]: ``[swa_pages + compress_pages, head_dim]`` bf16
|
||||
unified_kv[L]: ``[swa_pages + padded_compress_rows, head_dim]`` bf16
|
||||
- rows ``[0, swa_pages)`` = SWA ring (``req_pool_indices * swa_window + pos % swa_window``)
|
||||
- rows ``[swa_pages, ...)`` = compressed (``swa_pages + page_index``)
|
||||
"""
|
||||
@@ -403,6 +403,7 @@ class DeepSeekV4UnifiedKVPool:
|
||||
stage_ratios: List[int],
|
||||
num_slots: int,
|
||||
num_blocks: int,
|
||||
page_size: int,
|
||||
qk_nope_head_dim: int,
|
||||
qk_rope_head_dim: int,
|
||||
device: str,
|
||||
@@ -415,6 +416,7 @@ class DeepSeekV4UnifiedKVPool:
|
||||
self.num_slots = num_slots
|
||||
self.swa_pages = num_slots * self.swa_ring_size
|
||||
self.num_blocks = num_blocks
|
||||
self.page_size = page_size
|
||||
self.k_per_block = dict(self.K_PER_BLOCK)
|
||||
|
||||
bufs = []
|
||||
@@ -425,10 +427,14 @@ class DeepSeekV4UnifiedKVPool:
|
||||
else nullcontext()
|
||||
):
|
||||
for ratio in stage_ratios:
|
||||
compress_pages = self.num_blocks * self.k_per_block[ratio]
|
||||
# Pad by one extra page. The KV pool reserves a null slot
|
||||
# (token indices run 1..size).
|
||||
compress_rows = self.num_blocks * self.k_per_block[ratio]
|
||||
rows_per_page = self.page_size // ratio if ratio else 0
|
||||
padded_compress_rows = compress_rows + rows_per_page
|
||||
bufs.append(
|
||||
torch.zeros(
|
||||
self.swa_pages + compress_pages,
|
||||
self.swa_pages + padded_compress_rows,
|
||||
self.head_dim,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
@@ -579,6 +585,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
stage_ratios=stage_ratios,
|
||||
num_slots=self.num_req_slots,
|
||||
num_blocks=self.c128_size,
|
||||
page_size=page_size,
|
||||
qk_nope_head_dim=qk_nope_head_dim,
|
||||
qk_rope_head_dim=qk_rope_head_dim,
|
||||
device=device,
|
||||
@@ -645,6 +652,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
self._init_paged_compress_states(enable_memory_saver)
|
||||
|
||||
def get_unified_kv(self, layer_id: int) -> torch.Tensor:
|
||||
# Under HiCache the compressed region is loaded H->D per layer; wait for this
|
||||
# layer's transfer before attention reads it. No-op when HiCache is off.
|
||||
self.wait_layer_transfer(layer_id)
|
||||
return self.unified_kv_pool.get_unified_kv(layer_id - self._stage_start)
|
||||
|
||||
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor):
|
||||
@@ -665,7 +675,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
item_lens: List[int] = []
|
||||
|
||||
if self._unified_kv:
|
||||
# Unified buffer per layer: [swa_pages + compress_pages, head_dim].
|
||||
# Unified buffer per layer: [swa_pages + padded_compress_rows, head_dim].
|
||||
# Compressed region [swa_pages:] is page-contiguous (row swa_pages +
|
||||
# loc//ratio), so reuse the page-block PD transfer by offsetting the ptr
|
||||
# past the SWA ring and setting item_len = one page of rows. The SWA ring
|
||||
@@ -732,6 +742,43 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
item_lens.append(row_bytes)
|
||||
return data_ptrs, data_lens, item_lens
|
||||
|
||||
def unified_region_buffers(self, ratio: int) -> Tuple[List[torch.Tensor], int]:
|
||||
"""
|
||||
In unified_kv, swa/c4/c128 share one buffer with one slot per row. But the
|
||||
HiCache host pool transfers a whole page per indexed row, so we reshape the
|
||||
compressed region into the layout it expects: skip the SWA segment, reshape to
|
||||
one row per page, then cast to uint8.
|
||||
"""
|
||||
assert self._unified_kv, "unified_region_buffers requires unified_kv layout"
|
||||
assert ratio in (4, 128), f"unsupported compression ratio: {ratio}"
|
||||
|
||||
swa_pages = self.unified_kv_pool.swa_pages
|
||||
head_dim = self.unified_kv_pool.head_dim
|
||||
rows_per_page = self.page_size // ratio
|
||||
stage_ratios = self.compression_ratios[self._stage_start : self._stage_end]
|
||||
local_layer_ids = [i for i, r in enumerate(stage_ratios) if r == ratio]
|
||||
|
||||
views: List[torch.Tensor] = []
|
||||
for local_layer_id in local_layer_ids:
|
||||
buf = self.unified_kv_pool.kv_buffer[local_layer_id]
|
||||
compress_rows = buf.shape[0] - swa_pages
|
||||
assert compress_rows % rows_per_page == 0, (
|
||||
f"compressed rows {compress_rows} not a multiple of "
|
||||
f"rows_per_page {rows_per_page} for ratio {ratio}"
|
||||
)
|
||||
num_pages = compress_rows // rows_per_page
|
||||
page_view = (
|
||||
buf.narrow(0, swa_pages, compress_rows)
|
||||
.reshape(num_pages, rows_per_page * head_dim)
|
||||
.view(torch.uint8)
|
||||
)
|
||||
views.append(page_view)
|
||||
|
||||
item_bytes = (
|
||||
rows_per_page * head_dim * self.unified_kv_pool.kv_buffer[0].element_size()
|
||||
)
|
||||
return views, item_bytes
|
||||
|
||||
def get_state_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
|
||||
data_ptrs: List[int] = []
|
||||
data_lens: List[int] = []
|
||||
|
||||
@@ -270,6 +270,17 @@ def _deepseek_v4_num_host_pages(
|
||||
return full_host_pages, swa_host_pages
|
||||
|
||||
|
||||
def _dsv4_compressed_region_buffers(kvcache: Any, ratio: int) -> tuple[list, int]:
|
||||
"""
|
||||
Resolve ``(device_buffers, item_bytes)`` for a DeepSeek V4 C4/C128 main-KV
|
||||
HiCache pool, hiding the device KV layout from the stack builder.
|
||||
"""
|
||||
if getattr(kvcache, "_unified_kv", False):
|
||||
return kvcache.unified_region_buffers(ratio)
|
||||
pool = kvcache.c4_kv_pool if ratio == 4 else kvcache.c128_kv_pool
|
||||
return pool.kv_buffer, pool.bytes_per_page_padded
|
||||
|
||||
|
||||
def build_deepseek_v4_hicache_stack(
|
||||
*,
|
||||
params: CacheInitParams,
|
||||
@@ -291,13 +302,22 @@ def build_deepseek_v4_hicache_stack(
|
||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
|
||||
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
||||
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_num:
|
||||
raise ValueError(
|
||||
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
|
||||
f"got {len(kvcache.swa_kv_pool.kv_buffer)} buffers for "
|
||||
f"{transfer_layer_num} local layers"
|
||||
)
|
||||
swa_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
||||
|
||||
is_unified_kv = getattr(kvcache, "_unified_kv", False)
|
||||
if is_unified_kv:
|
||||
# unified_kv keeps the SWA ring inside the unified pool and never offloads it,
|
||||
# so there is no separate SWA host pool to map.
|
||||
swa_layer_mapping = {}
|
||||
else:
|
||||
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_num:
|
||||
raise ValueError(
|
||||
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
|
||||
f"got {len(kvcache.swa_kv_pool.kv_buffer)} buffers for "
|
||||
f"{transfer_layer_num} local layers"
|
||||
)
|
||||
swa_layer_mapping = {
|
||||
layer_id: layer_id for layer_id in range(transfer_layer_num)
|
||||
}
|
||||
|
||||
c4_layer_mapping = {}
|
||||
c128_layer_mapping = {}
|
||||
@@ -328,16 +348,6 @@ def build_deepseek_v4_hicache_stack(
|
||||
logical_host_pool = LogicalHostPool(
|
||||
num_host_pages * page_size, page_size, layout=server_args.hicache_mem_layout
|
||||
)
|
||||
swa_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.SWA),
|
||||
device_buffers=kvcache.swa_kv_pool.kv_buffer,
|
||||
item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded,
|
||||
num_host_pages=swa_num_host_pages,
|
||||
slot_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
||||
entries = [
|
||||
build_pool_entry(
|
||||
name=PoolName.KV,
|
||||
@@ -347,24 +357,39 @@ def build_deepseek_v4_hicache_stack(
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
is_anchor=True,
|
||||
),
|
||||
build_pool_entry(
|
||||
name=PoolName.SWA,
|
||||
host_pool=swa_host_pool,
|
||||
device_pool=kvcache.swa_kv_pool,
|
||||
layer_mapping=swa_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
host_evict_fn=host_swa_evict_fn,
|
||||
device_evict_fn=device_swa_evict_fn,
|
||||
device_alloc_fn=swa_attn_allocator.alloc,
|
||||
device_free_fn=swa_attn_allocator.free,
|
||||
),
|
||||
]
|
||||
|
||||
if not is_unified_kv:
|
||||
swa_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.SWA),
|
||||
device_buffers=kvcache.swa_kv_pool.kv_buffer,
|
||||
item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded,
|
||||
num_host_pages=swa_num_host_pages,
|
||||
slot_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
||||
entries.append(
|
||||
build_pool_entry(
|
||||
name=PoolName.SWA,
|
||||
host_pool=swa_host_pool,
|
||||
device_pool=kvcache.swa_kv_pool,
|
||||
layer_mapping=swa_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
host_evict_fn=host_swa_evict_fn,
|
||||
device_evict_fn=device_swa_evict_fn,
|
||||
device_alloc_fn=swa_attn_allocator.alloc,
|
||||
device_free_fn=swa_attn_allocator.free,
|
||||
)
|
||||
)
|
||||
|
||||
if c4_layer_mapping:
|
||||
c4_device_buffers, c4_item_bytes = _dsv4_compressed_region_buffers(kvcache, 4)
|
||||
c4_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4),
|
||||
device_buffers=kvcache.c4_kv_pool.kv_buffer,
|
||||
item_bytes=kvcache.c4_kv_pool.bytes_per_page_padded,
|
||||
device_buffers=c4_device_buffers,
|
||||
item_bytes=c4_item_bytes,
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
@@ -382,28 +407,6 @@ def build_deepseek_v4_hicache_stack(
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
c4_state_host_pool = DeepSeekV4StateHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_STATE),
|
||||
state_pools=[
|
||||
kvcache.compress_state_pools[layer_id]
|
||||
for layer_id in c4_state_global_layers
|
||||
],
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE),
|
||||
state_pools=[
|
||||
kvcache.indexer_compress_state_pools[layer_id]
|
||||
for layer_id in c4_state_global_layers
|
||||
],
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
entries.extend(
|
||||
[
|
||||
build_pool_entry(
|
||||
@@ -420,28 +423,59 @@ def build_deepseek_v4_hicache_stack(
|
||||
layer_mapping=c4_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4_STATE,
|
||||
host_pool=c4_state_host_pool,
|
||||
device_pool=None,
|
||||
layer_mapping=c4_state_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
|
||||
host_pool=c4_indexer_state_host_pool,
|
||||
device_pool=None,
|
||||
layer_mapping=c4_state_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
if not is_unified_kv:
|
||||
c4_state_host_pool = DeepSeekV4StateHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_STATE),
|
||||
state_pools=[
|
||||
kvcache.compress_state_pools[layer_id]
|
||||
for layer_id in c4_state_global_layers
|
||||
],
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE),
|
||||
state_pools=[
|
||||
kvcache.indexer_compress_state_pools[layer_id]
|
||||
for layer_id in c4_state_global_layers
|
||||
],
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
)
|
||||
entries.extend(
|
||||
[
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4_STATE,
|
||||
host_pool=c4_state_host_pool,
|
||||
device_pool=None,
|
||||
layer_mapping=c4_state_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
|
||||
host_pool=c4_indexer_state_host_pool,
|
||||
device_pool=None,
|
||||
layer_mapping=c4_state_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
if c128_layer_mapping:
|
||||
c128_device_buffers, c128_item_bytes = _dsv4_compressed_region_buffers(
|
||||
kvcache, 128
|
||||
)
|
||||
c128_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C128),
|
||||
device_buffers=kvcache.c128_kv_pool.kv_buffer,
|
||||
item_bytes=kvcache.c128_kv_pool.bytes_per_page_padded,
|
||||
device_buffers=c128_device_buffers,
|
||||
item_bytes=c128_item_bytes,
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
@@ -744,13 +778,18 @@ class _DeepSeekV4Strategy(StackStrategy):
|
||||
)
|
||||
if name in host_pool_group.entry_map
|
||||
]
|
||||
component_host_pools = {
|
||||
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
||||
}
|
||||
if PoolName.SWA in host_pool_group.entry_map:
|
||||
component_host_pools[ComponentType.SWA] = host_pool_group.get_pool(
|
||||
PoolName.SWA
|
||||
)
|
||||
|
||||
return StackBuildResult(
|
||||
host_pool_group=host_pool_group,
|
||||
cache_controller=cache_controller,
|
||||
component_host_pools={
|
||||
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
||||
ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA),
|
||||
},
|
||||
component_host_pools=component_host_pools,
|
||||
sidecars=sidecars,
|
||||
transfer_layer_num=kvcache.end_layer - kvcache.start_layer,
|
||||
pools_desc="KV + SWA + DeepSeekV4 sidecars",
|
||||
|
||||
@@ -121,12 +121,20 @@ class SWAComponent(TreeComponent):
|
||||
ct = self.component_type
|
||||
state = {"len": float("inf")}
|
||||
|
||||
# unified_kv never caches the SWA ring (per-request, not content-stable),
|
||||
# so SWA bookkeeping must not gate the match here.
|
||||
swa_device_only_hicache = (
|
||||
self._swa_kv_pool_host is None and self.cache.cache_controller is not None
|
||||
)
|
||||
|
||||
def validator(node: UnifiedTreeNode) -> bool:
|
||||
cd = node.component_data[ct]
|
||||
# HiCache: a host-only tombstone is a valid match boundary too
|
||||
# — load_back will restore SWA from host before use.
|
||||
if cd.value is None and (match_device_only or cd.host_value is None):
|
||||
state["len"] = 0
|
||||
if swa_device_only_hicache and (node.backuped or not node.evicted):
|
||||
return True
|
||||
return False
|
||||
state["len"] += len(node.key)
|
||||
return state["len"] >= sliding_window_size
|
||||
@@ -612,6 +620,10 @@ class SWAComponent(TreeComponent):
|
||||
) -> Optional[list[PoolTransfer]]:
|
||||
ct = self.component_type
|
||||
|
||||
# unified_kv keeps SWA as a device-only ring.
|
||||
if self._swa_kv_pool_host is None and self.cache.cache_controller is not None:
|
||||
return None
|
||||
|
||||
if phase == CacheTransferPhase.BACKUP_HOST:
|
||||
cd = node.component_data[ct]
|
||||
if cd.value is None:
|
||||
|
||||
@@ -2484,6 +2484,23 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
swa = self.components.get(ComponentType.SWA)
|
||||
return swa.sliding_window_size if swa else None
|
||||
|
||||
def swa_reprefill_tail_tokens(self) -> int:
|
||||
"""
|
||||
Only unified_kv + HiCache needs this: SWA lives in a per-request ring
|
||||
(state_slot/pos), not content-stable and never offloaded to host, so a
|
||||
reused prefix's trailing sliding window would read another request's
|
||||
stale ring slots. Re-prefilling that window rewrites this request's ring
|
||||
(what plain radix reuse does via its SWA match gate). 0 for every other
|
||||
layout.
|
||||
"""
|
||||
swa = self.components.get(ComponentType.SWA)
|
||||
unified_compress_only_hicache = (
|
||||
self.cache_controller is not None
|
||||
and swa is not None
|
||||
and swa._swa_kv_pool_host is None
|
||||
)
|
||||
return swa.sliding_window_size if unified_compress_only_hicache else 0
|
||||
|
||||
def supports_swa(self) -> bool:
|
||||
return ComponentType.SWA in self.components
|
||||
|
||||
|
||||
@@ -5653,26 +5653,6 @@ class ServerArgs:
|
||||
# Step 2: Storage-layout normalization without changing io backend.
|
||||
self._resolve_storage_layout_compatibility()
|
||||
|
||||
# Step 3: HiCache is not yet supported with the DeepSeek-V4 hip unified_kv
|
||||
# layout, so fall back to the default tilelang FlashMLA backend.
|
||||
self._resolve_unified_kv_hicache_compatibility()
|
||||
|
||||
def _resolve_unified_kv_hicache_compatibility(self):
|
||||
# The DeepSeek-V4 unified_kv layout (SGLANG_HACK_FLASHMLA_BACKEND=
|
||||
# unified_kv_triton) keeps swa/c4/c128 in a single per-layer buffer and
|
||||
# has no HiCache host-pool support yet, so reset the backend to the
|
||||
# default (tilelang) so the server still starts.
|
||||
if not self.enable_hierarchical_cache:
|
||||
return
|
||||
|
||||
if envs.SGLANG_HACK_FLASHMLA_BACKEND.get() == "unified_kv_triton":
|
||||
envs.SGLANG_HACK_FLASHMLA_BACKEND.set("tilelang")
|
||||
logger.warning(
|
||||
"SGLANG_HACK_FLASHMLA_BACKEND=unified_kv_triton is not yet "
|
||||
"compatible with --enable-hierarchical-cache; falling back to "
|
||||
"SGLANG_HACK_FLASHMLA_BACKEND=tilelang."
|
||||
)
|
||||
|
||||
def _resolve_layout_io_compatibility(self):
|
||||
if (
|
||||
self.hicache_mem_layout == "page_first_direct"
|
||||
|
||||
Reference in New Issue
Block a user