Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 53 additions & 34 deletions python/sglang/srt/mem_cache/kv_cache_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,10 @@
PageMajorMHATokenToKVPool,
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.platforms import current_platform
from sglang.srt.runtime_context import (
Expand Down Expand Up @@ -482,35 +486,43 @@ def _init_pools(
# pool must be sized by that space.
draft_virtual_id_space: Optional[int] = None
if self.is_draft_worker and token_to_kv_pool_allocator is not None:
from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator,
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
token_to_kv_pool_allocator,
(
UnifiedMambaTokenToKVPoolAllocator,
UnifiedSWATokenToKVPoolAllocator,
),
):
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, (
"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}"
)
# 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)
draft_virtual_id_space = (
(draft_virtual_id_space + page - 1) // page * page
)
sizes = msgspec.structs.replace(
sizes, max_total_num_tokens=draft_virtual_id_space
)
size_overrides = {
"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
if req_to_token_pool is None:
Expand Down Expand Up @@ -558,7 +570,8 @@ def _init_pools(
assert token_to_kv_pool.size >= draft_virtual_id_space, (
"draft token_to_kv_pool smaller than the shared unified "
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 "
"bounds."
)
Expand Down Expand Up @@ -1960,27 +1973,33 @@ def _build_token_to_kv_pool_allocator(
else:
assert self.is_draft_worker
if self.is_hybrid_swa:
if self.draft_swa_full_capacity:
# Banded depth: the SWA ring is full draft capacity, so use
# an IDENTITY full->swa mapping — store and read locs both
# equal out_cache_loc, and a slot is never evicted before
# the request frees it. The window itself is enforced by the
# FA sliding-window kernel, not by the ring. Layout mirrors
# SWATokenToKVPoolAllocator's mapping (size + page_size
# entries + trailing -1 sentinel so a -1 last_loc maps
# to -1).
if isinstance(
token_to_kv_pool_allocator,
DeepSeekV4HiSparseTokenToKVPoolAllocator,
):
swa_allocator = token_to_kv_pool_allocator.logical_attn_allocator
else:
swa_allocator = token_to_kv_pool_allocator
uses_unified_virtual_ids = isinstance(
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
identity_mapping = torch.arange(
n + 1, dtype=torch.int64, device=self.device
)
identity_mapping[-1] = -1
token_to_kv_pool.register_mapping(identity_mapping)
else:
swa_allocator = getattr(
token_to_kv_pool_allocator,
"logical_attn_allocator",
token_to_kv_pool_allocator,
)
elif not uses_unified_virtual_ids:
assert isinstance(swa_allocator, SWATokenToKVPoolAllocator)
token_to_kv_pool.register_mapping(
swa_allocator.full_to_swa_index_mapping
Expand Down
8 changes: 8 additions & 0 deletions python/sglang/srt/mem_cache/multi_ended_allocator.py
Original file line number Diff line number Diff line change
Expand Up @@ -2890,6 +2890,10 @@ def size_full(self) -> int:
# Widened like `size`: a logical token capacity, not a row count.
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
def size_mamba(self) -> int:
return self.mamba_allocator.max_slots - 1
Expand Down Expand Up @@ -3390,6 +3394,10 @@ def _ask_float_for_room(self, need_tokens: int) -> None:
# (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.

@property
def draft_virtual_id_space(self) -> int:
return self.full_attn_allocator.max_slots - 1

def debug_print(self) -> str:
return (
f"#full-available={self.full_attn_allocator.available_size()}, "
Expand Down
80 changes: 72 additions & 8 deletions python/sglang/srt/model_executor/pool_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from __future__ import annotations

import logging
from bisect import bisect_right
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional

Expand Down Expand Up @@ -465,6 +466,7 @@ def __init__(self, kvc: KVCacheConfigurator):
self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
self._sliding_window_size = kvc.sliding_window_size
self._page_size = kvc.page_size
self._enable_unified_memory = get_memory().enable_unified_memory

if model_config.attention_arch == AttentionArch.MLA:
# MLA pool sizing uses latent dimensions rather than MHA heads.
Expand Down Expand Up @@ -551,6 +553,9 @@ def __init__(self, kvc: KVCacheConfigurator):

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.
#
# Hybrid (full_layers > 0): max_total = full_tokens, so cell_size accounts
Expand Down Expand Up @@ -579,6 +584,50 @@ def __init__(self, kvc: KVCacheConfigurator):
+ 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(
self, max_total_num_tokens: int, page_size: int
) -> MemoryPoolConfig:
Expand Down Expand Up @@ -630,7 +679,16 @@ def align_page_size(x: int) -> int:
def calculate_pool_sizes(
self, available_bytes: int, page_size: int
) -> MemoryPoolConfig:
max_total_num_tokens = int(available_bytes // self._cell_size)
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)
return self._solve_pool_sizes(max_total_num_tokens, page_size)

def calculate_pool_sizes_from_max_tokens(
Expand Down Expand Up @@ -716,13 +774,19 @@ def calculate_pool_sizes(
* self._swa_per_token
* (self._swa_layers_num + self._draft_swa_layers_num)
)
full_cell_size = (
self._full_per_token * (self._full_layers_num + self._draft_full_layers_num)
+ self._swa_per_token * self._draft_swa_full_layers_num
)
full_tokens = (
int((available_bytes - fixed_swa_bytes) // full_cell_size) // page_size
) * page_size
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 = (
self._full_per_token
* (self._full_layers_num + self._draft_full_layers_num)
+ self._swa_per_token * self._draft_swa_full_layers_num
)
full_tokens = (
int((available_bytes - fixed_swa_bytes) // full_cell_size) // page_size
) * page_size
if full_tokens <= 0:
raise RuntimeError(
f"SWA pool cap ({swa_tokens} tokens, "
Expand Down
Loading