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:
Yonghao Zhuang
2026-09-03 10:44:57 -07:00
committed by GitHub
co-authored by yhzhuang Lianmin Zheng
parent 33a22b1b08
commit 23ab10a63e
3 changed files with 133 additions and 42 deletions
@@ -14,6 +14,7 @@ Two entry points, same core computation:
from __future__ import annotations
import logging
from bisect import bisect_right
from dataclasses import dataclass
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._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.
@@ -555,6 +557,9 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
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
@@ -583,6 +588,50 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
+ 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:
@@ -634,7 +683,16 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
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(
@@ -720,13 +778,19 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
* 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, "