Support speculative decoding with unified SWA memory (#36403)
Co-authored-by: yhzhuang <yhzhuang@fb.com> Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
co-authored by
yhzhuang
Lianmin Zheng
parent
33a22b1b08
commit
23ab10a63e
@@ -64,6 +64,10 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
PageMajorMHATokenToKVPool,
|
PageMajorMHATokenToKVPool,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.multi_ended_allocator import (
|
||||||
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
@@ -500,35 +504,43 @@ class KVCacheConfigurator:
|
|||||||
# pool must be sized by that space.
|
# pool must be sized by that space.
|
||||||
draft_virtual_id_space: Optional[int] = None
|
draft_virtual_id_space: Optional[int] = None
|
||||||
if self.is_draft_worker and token_to_kv_pool_allocator is not None:
|
if self.is_draft_worker and token_to_kv_pool_allocator is not None:
|
||||||
from sglang.srt.mem_cache.multi_ended_allocator import (
|
if isinstance(
|
||||||
|
token_to_kv_pool_allocator,
|
||||||
|
(
|
||||||
UnifiedMambaTokenToKVPoolAllocator,
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
)
|
),
|
||||||
|
|
||||||
if isinstance(token_to_kv_pool_allocator, UnifiedSWATokenToKVPoolAllocator):
|
|
||||||
raise ValueError(
|
|
||||||
"Speculative decoding with --enable-unified-memory is only "
|
|
||||||
"supported for hybrid-Mamba targets; the unified hybrid-SWA "
|
|
||||||
"pool's draft sizing (virtual-id space) is not wired yet."
|
|
||||||
)
|
|
||||||
if isinstance(
|
|
||||||
token_to_kv_pool_allocator, UnifiedMambaTokenToKVPoolAllocator
|
|
||||||
):
|
):
|
||||||
draft_virtual_id_space = token_to_kv_pool_allocator.size_full
|
draft_virtual_id_space = (
|
||||||
|
token_to_kv_pool_allocator.draft_virtual_id_space
|
||||||
|
)
|
||||||
assert draft_virtual_id_space >= sizes.max_total_num_tokens, (
|
assert draft_virtual_id_space >= sizes.max_total_num_tokens, (
|
||||||
"unified allocator virtual space smaller than the token "
|
"unified allocator virtual space smaller than the token "
|
||||||
f"budget: size_full={draft_virtual_id_space} < "
|
f"budget: virtual_id_space={draft_virtual_id_space} < "
|
||||||
f"max_total_num_tokens={sizes.max_total_num_tokens}"
|
f"max_total_num_tokens={sizes.max_total_num_tokens}"
|
||||||
)
|
)
|
||||||
# Round UP to page alignment (paged draft backends view the
|
# Round UP to page alignment (paged draft backends view the
|
||||||
# pool as (-1, page_size, H, D); size_full is not aligned).
|
# pool as (-1, page_size, H, D); the virtual space is not aligned).
|
||||||
page = max(int(self.pool_page_size or 1), 1)
|
page = max(int(self.pool_page_size or 1), 1)
|
||||||
draft_virtual_id_space = (
|
draft_virtual_id_space = (
|
||||||
(draft_virtual_id_space + page - 1) // page * page
|
(draft_virtual_id_space + page - 1) // page * page
|
||||||
)
|
)
|
||||||
sizes = msgspec.structs.replace(
|
size_overrides = {
|
||||||
sizes, max_total_num_tokens=draft_virtual_id_space
|
"max_total_num_tokens": draft_virtual_id_space,
|
||||||
|
}
|
||||||
|
if (
|
||||||
|
isinstance(
|
||||||
|
token_to_kv_pool_allocator,
|
||||||
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
)
|
)
|
||||||
|
and self.is_hybrid_swa
|
||||||
|
):
|
||||||
|
size_overrides["full_max_total_num_tokens"] = draft_virtual_id_space
|
||||||
|
if not self.is_hybrid_swa_mtp_draft or self.draft_swa_full_capacity:
|
||||||
|
size_overrides["swa_max_total_num_tokens"] = (
|
||||||
|
draft_virtual_id_space
|
||||||
|
)
|
||||||
|
sizes = msgspec.structs.replace(sizes, **size_overrides)
|
||||||
|
|
||||||
# Initialize req_to_token_pool
|
# Initialize req_to_token_pool
|
||||||
if req_to_token_pool is None:
|
if req_to_token_pool is None:
|
||||||
@@ -576,7 +588,8 @@ class KVCacheConfigurator:
|
|||||||
assert token_to_kv_pool.size >= draft_virtual_id_space, (
|
assert token_to_kv_pool.size >= draft_virtual_id_space, (
|
||||||
"draft token_to_kv_pool smaller than the shared unified "
|
"draft token_to_kv_pool smaller than the shared unified "
|
||||||
f"allocator's virtual-id space: pool size="
|
f"allocator's virtual-id space: pool size="
|
||||||
f"{token_to_kv_pool.size} < size_full={draft_virtual_id_space}; "
|
f"{token_to_kv_pool.size} < "
|
||||||
|
f"virtual_id_space={draft_virtual_id_space}; "
|
||||||
"verify-window writes at high virtual ids would go out of "
|
"verify-window writes at high virtual ids would go out of "
|
||||||
"bounds."
|
"bounds."
|
||||||
)
|
)
|
||||||
@@ -1979,27 +1992,33 @@ class KVCacheConfigurator:
|
|||||||
else:
|
else:
|
||||||
assert self.is_draft_worker
|
assert self.is_draft_worker
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
if self.draft_swa_full_capacity:
|
if isinstance(
|
||||||
# Banded depth: the SWA ring is full draft capacity, so use
|
token_to_kv_pool_allocator,
|
||||||
# an IDENTITY full->swa mapping — store and read locs both
|
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
||||||
# equal out_cache_loc, and a slot is never evicted before
|
):
|
||||||
# the request frees it. The window itself is enforced by the
|
swa_allocator = token_to_kv_pool_allocator.logical_attn_allocator
|
||||||
# FA sliding-window kernel, not by the ring. Layout mirrors
|
else:
|
||||||
# SWATokenToKVPoolAllocator's mapping (size + page_size
|
swa_allocator = token_to_kv_pool_allocator
|
||||||
# entries + trailing -1 sentinel so a -1 last_loc maps
|
uses_unified_virtual_ids = isinstance(
|
||||||
# to -1).
|
swa_allocator, UnifiedSWATokenToKVPoolAllocator
|
||||||
|
)
|
||||||
|
has_draft_swa_layers = (
|
||||||
|
not self.is_hybrid_swa_mtp_draft or self.draft_swa_full_capacity
|
||||||
|
)
|
||||||
|
if self.draft_swa_full_capacity or (
|
||||||
|
uses_unified_virtual_ids and has_draft_swa_layers
|
||||||
|
):
|
||||||
|
# The draft pool owns independent KV but consumes the target
|
||||||
|
# allocator's virtual ids directly. Size its SWA side for that
|
||||||
|
# whole space and use an identity mapping. The trailing -1
|
||||||
|
# sentinel keeps a -1 last_loc mapped to -1.
|
||||||
n = sizes.full_max_total_num_tokens + self.page_size
|
n = sizes.full_max_total_num_tokens + self.page_size
|
||||||
identity_mapping = torch.arange(
|
identity_mapping = torch.arange(
|
||||||
n + 1, dtype=torch.int64, device=self.device
|
n + 1, dtype=torch.int64, device=self.device
|
||||||
)
|
)
|
||||||
identity_mapping[-1] = -1
|
identity_mapping[-1] = -1
|
||||||
token_to_kv_pool.register_mapping(identity_mapping)
|
token_to_kv_pool.register_mapping(identity_mapping)
|
||||||
else:
|
elif not uses_unified_virtual_ids:
|
||||||
swa_allocator = getattr(
|
|
||||||
token_to_kv_pool_allocator,
|
|
||||||
"logical_attn_allocator",
|
|
||||||
token_to_kv_pool_allocator,
|
|
||||||
)
|
|
||||||
assert isinstance(swa_allocator, SWATokenToKVPoolAllocator)
|
assert isinstance(swa_allocator, SWATokenToKVPoolAllocator)
|
||||||
token_to_kv_pool.register_mapping(
|
token_to_kv_pool.register_mapping(
|
||||||
swa_allocator.full_to_swa_index_mapping
|
swa_allocator.full_to_swa_index_mapping
|
||||||
|
|||||||
@@ -2898,6 +2898,10 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
# Widened like `size`: a logical token capacity, not a row count.
|
# Widened like `size`: a logical token capacity, not a row count.
|
||||||
return (self.full_attn_allocator.max_slots - 1) * get_parallel().attn_dcp_size
|
return (self.full_attn_allocator.max_slots - 1) * get_parallel().attn_dcp_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def draft_virtual_id_space(self) -> int:
|
||||||
|
return self.size_full
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def size_mamba(self) -> int:
|
def size_mamba(self) -> int:
|
||||||
return self.mamba_allocator.max_slots - 1
|
return self.mamba_allocator.max_slots - 1
|
||||||
@@ -3401,6 +3405,10 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
# (set to the static caps). We do NOT report `max_slots - 1`: under unified
|
# (set to the static caps). We do NOT report `max_slots - 1`: under unified
|
||||||
# memory pool that ~= full_max + swa_max and would over-promise.
|
# memory pool that ~= full_max + swa_max and would over-promise.
|
||||||
|
|
||||||
|
@property
|
||||||
|
def draft_virtual_id_space(self) -> int:
|
||||||
|
return self.full_attn_allocator.max_slots - 1
|
||||||
|
|
||||||
def debug_print(self) -> str:
|
def debug_print(self) -> str:
|
||||||
return (
|
return (
|
||||||
f"#full-available={self.full_attn_allocator.available_size()}, "
|
f"#full-available={self.full_attn_allocator.available_size()}, "
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ Two entry points, same core computation:
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from bisect import bisect_right
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
@@ -469,6 +470,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
|
self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
|
||||||
self._sliding_window_size = kvc.sliding_window_size
|
self._sliding_window_size = kvc.sliding_window_size
|
||||||
self._page_size = kvc.page_size
|
self._page_size = kvc.page_size
|
||||||
|
self._enable_unified_memory = get_memory().enable_unified_memory
|
||||||
|
|
||||||
if model_config.attention_arch == AttentionArch.MLA:
|
if model_config.attention_arch == AttentionArch.MLA:
|
||||||
# MLA pool sizing uses latent dimensions rather than MHA heads.
|
# MLA pool sizing uses latent dimensions rather than MHA heads.
|
||||||
@@ -555,6 +557,9 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
|
|
||||||
self._draft_cell_size = _dflash_draft_cell_size(kvc)
|
self._draft_cell_size = _dflash_draft_cell_size(kvc)
|
||||||
|
|
||||||
|
self._recompute_cell_size()
|
||||||
|
|
||||||
|
def _recompute_cell_size(self) -> None:
|
||||||
# Bytes per token of max_total_num_tokens.
|
# Bytes per token of max_total_num_tokens.
|
||||||
#
|
#
|
||||||
# Hybrid (full_layers > 0): max_total = full_tokens, so cell_size accounts
|
# Hybrid (full_layers > 0): max_total = full_tokens, so cell_size accounts
|
||||||
@@ -583,6 +588,50 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
+ self._draft_cell_size
|
+ self._draft_cell_size
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _draft_pool_bytes_per_token(self) -> int:
|
||||||
|
return int(
|
||||||
|
self._full_per_token * self._draft_full_layers_num
|
||||||
|
+ self._swa_per_token
|
||||||
|
* (self._draft_swa_layers_num + self._draft_swa_full_layers_num)
|
||||||
|
+ self._draft_cell_size
|
||||||
|
)
|
||||||
|
|
||||||
|
def _max_unified_full_tokens(
|
||||||
|
self,
|
||||||
|
available_bytes: int,
|
||||||
|
page_size: int,
|
||||||
|
fixed_swa_tokens: Optional[int] = None,
|
||||||
|
) -> int:
|
||||||
|
"""Find the largest page-aligned full capacity whose allocations fit."""
|
||||||
|
draft_bytes_per_token = self._draft_pool_bytes_per_token()
|
||||||
|
target_full_bytes_per_token = self._full_per_token * self._full_layers_num
|
||||||
|
target_swa_bytes_per_token = self._swa_per_token * self._swa_layers_num
|
||||||
|
assert target_full_bytes_per_token > 0
|
||||||
|
|
||||||
|
def allocation_bytes(full_pages: int) -> int:
|
||||||
|
full_tokens = full_pages * page_size
|
||||||
|
swa_tokens = (
|
||||||
|
fixed_swa_tokens
|
||||||
|
if fixed_swa_tokens is not None
|
||||||
|
else int(full_tokens * self._swa_full_tokens_ratio)
|
||||||
|
// page_size
|
||||||
|
* page_size
|
||||||
|
)
|
||||||
|
target_bytes = (
|
||||||
|
full_tokens * target_full_bytes_per_token
|
||||||
|
+ swa_tokens * target_swa_bytes_per_token
|
||||||
|
)
|
||||||
|
virtual_span = max(target_bytes // target_full_bytes_per_token - 1, 0)
|
||||||
|
draft_tokens = ceil_align(virtual_span, page_size) + page_size
|
||||||
|
return target_bytes + draft_tokens * draft_bytes_per_token
|
||||||
|
|
||||||
|
max_pages = available_bytes // target_full_bytes_per_token // page_size
|
||||||
|
full_pages = (
|
||||||
|
bisect_right(range(max_pages + 1), available_bytes, key=allocation_bytes)
|
||||||
|
- 1
|
||||||
|
)
|
||||||
|
return max(full_pages, 0) * page_size
|
||||||
|
|
||||||
def _solve_pool_sizes(
|
def _solve_pool_sizes(
|
||||||
self, max_total_num_tokens: int, page_size: int
|
self, max_total_num_tokens: int, page_size: int
|
||||||
) -> MemoryPoolConfig:
|
) -> MemoryPoolConfig:
|
||||||
@@ -634,6 +683,15 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
def calculate_pool_sizes(
|
def calculate_pool_sizes(
|
||||||
self, available_bytes: int, page_size: int
|
self, available_bytes: int, page_size: int
|
||||||
) -> MemoryPoolConfig:
|
) -> MemoryPoolConfig:
|
||||||
|
if (
|
||||||
|
self._enable_unified_memory
|
||||||
|
and self._full_layers_num > 0
|
||||||
|
and self._draft_pool_bytes_per_token() > 0
|
||||||
|
):
|
||||||
|
max_total_num_tokens = self._max_unified_full_tokens(
|
||||||
|
available_bytes, page_size
|
||||||
|
)
|
||||||
|
else:
|
||||||
max_total_num_tokens = int(available_bytes // self._cell_size)
|
max_total_num_tokens = int(available_bytes // self._cell_size)
|
||||||
return self._solve_pool_sizes(max_total_num_tokens, page_size)
|
return self._solve_pool_sizes(max_total_num_tokens, page_size)
|
||||||
|
|
||||||
@@ -720,8 +778,14 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
* self._swa_per_token
|
* self._swa_per_token
|
||||||
* (self._swa_layers_num + self._draft_swa_layers_num)
|
* (self._swa_layers_num + self._draft_swa_layers_num)
|
||||||
)
|
)
|
||||||
|
if self._enable_unified_memory and self._draft_pool_bytes_per_token() > 0:
|
||||||
|
full_tokens = self._max_unified_full_tokens(
|
||||||
|
available_bytes, page_size, fixed_swa_tokens=swa_tokens
|
||||||
|
)
|
||||||
|
else:
|
||||||
full_cell_size = (
|
full_cell_size = (
|
||||||
self._full_per_token * (self._full_layers_num + self._draft_full_layers_num)
|
self._full_per_token
|
||||||
|
* (self._full_layers_num + self._draft_full_layers_num)
|
||||||
+ self._swa_per_token * self._draft_swa_full_layers_num
|
+ self._swa_per_token * self._draft_swa_full_layers_num
|
||||||
)
|
)
|
||||||
full_tokens = (
|
full_tokens = (
|
||||||
|
|||||||
Reference in New Issue
Block a user