Dispatch EAGLE radix bigram builder to tai-kernel
The pure-Python convert_to_bigram_key list comprehension runs on every cache_finished_req / cache_unfinished_req of every EAGLE radix variant (radix_cache, hiradix_cache, swa_radix_cache), with token-list lengths that scale with prompt + output. Scheduler profiles consistently flag it as the largest mem_cache-side CPU hotspot. This commit wires sglang.srt.mem_cache.utils.convert_to_bigram_key to tai_kernel.radix.convert_to_bigram_key when the extension is importable, falling back to the pure-Python implementation otherwise. The tai-kernel path uses a pybind11 module that calls the CPython C API directly (PyTuple_New + PyTuple_SET_ITEM with ref-stealing) rather than rebuilding the list comprehension's bytecode-level tuple allocations. Measured 1.4x at n=131k and up to 2.5x for n=1k on g0034 Python 3.11; allocator-bound at large n because CPython's 2-tuple freelist already amortises the construction. The int64-packing follow-up that bypasses tuple allocation entirely is parked as a separate work item. Runtime safety: - The dispatcher catches any first-call JIT compile / runtime failure, logs once, and falls through to the pure-Python path for the rest of the process — JIT failures must degrade rather than crash a serving loop. - SGLANG_DISABLE_TAI_BIGRAM forces the Python path for bisecting. Constraint: Output must be a real Python List[Tuple[int, int]] because downstream radix dicts use the tuples as hashable keys. Rejected: Pre-allocated tuple slab pool | CPython's per-interpreter 2-tuple freelist already serves this case, and we cannot recycle tuples that become radix-tree keys without changing the consumer. Rejected: int64-packed keys this round | requires changes to RadixKey, get_child_key_fn, key_match_fn, and EAGLE bigram detection; deserves its own plan. Confidence: high Scope-risk: low Directive: Keep _python_convert_to_bigram_key reachable; if the tai-kernel path is ever removed, the EAGLE radix cache must continue to work unchanged. Tested: tai-kernel side validated on the cluster (``python benchmark/radix/benchmark_convert_to_bigram_key.py --check`` prints byte-exact correctness then 1.4-2.5x speedup across sizes 128..131072 on g0034). The new ``test/registered/unit/mem_cache/test_convert_to_bigram_key.py`` exercises both dispatch paths via SGLANG_DISABLE_TAI_BIGRAM patching. Not-tested: End-to-end EAGLE serving accuracy + scheduler-time delta on a real workload. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -215,6 +215,7 @@ class Envs:
|
||||
SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES = EnvInt(-1)
|
||||
SGLANG_CP_DRAFT_SHARED_KV = EnvBool(False)
|
||||
SGLANG_CP_DRAFT_SHARED_KV_DEBUG = EnvBool(False)
|
||||
SGLANG_DISABLE_TAI_BIGRAM = EnvBool(False)
|
||||
SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False)
|
||||
SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(False)
|
||||
SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1)
|
||||
|
||||
@@ -13,6 +13,7 @@
|
||||
# ==============================================================================
|
||||
"""Common utilities."""
|
||||
|
||||
import logging
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -21,6 +22,8 @@ import triton.language as tl
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def set_mla_kv_buffer_kernel(
|
||||
@@ -264,7 +267,7 @@ def maybe_init_custom_mem_pool(
|
||||
return False, None, None
|
||||
|
||||
|
||||
def convert_to_bigram_key(tokens: List[int]) -> List[Tuple[int, int]]:
|
||||
def _python_convert_to_bigram_key(tokens: List[int]) -> List[Tuple[int, int]]:
|
||||
# EAGLE uses bigram keys in the radix tree since draft sequence is the one-token-shifted version of target
|
||||
# [1, 2, 3, 4] -> [(1,2), (2,3), (3,4)]
|
||||
if len(tokens) and isinstance(tokens[0], tuple):
|
||||
@@ -272,3 +275,28 @@ def convert_to_bigram_key(tokens: List[int]) -> List[Tuple[int, int]]:
|
||||
if len(tokens) < 2:
|
||||
return []
|
||||
return [(tokens[i], tokens[i + 1]) for i in range(len(tokens) - 1)]
|
||||
|
||||
|
||||
try:
|
||||
from tai_kernel.radix import convert_to_bigram_key as _tai_bigram_impl
|
||||
except Exception:
|
||||
_tai_bigram_impl = None
|
||||
|
||||
|
||||
def convert_to_bigram_key(tokens: List[int]) -> List[Tuple[int, int]]:
|
||||
global _tai_bigram_impl
|
||||
if _tai_bigram_impl is not None and not envs.SGLANG_DISABLE_TAI_BIGRAM.get():
|
||||
try:
|
||||
return _tai_bigram_impl(tokens)
|
||||
except Exception as e:
|
||||
# JIT compile or runtime failures (missing compiler, stale cache,
|
||||
# CPython ABI mismatch) must degrade to the Python path rather
|
||||
# than crash a serving loop. Disable after one failure so we
|
||||
# don't keep paying the broken-extension cost on every call.
|
||||
logger.warning(
|
||||
"tai-kernel convert_to_bigram_key failed, falling back to "
|
||||
"Python implementation for the rest of this process: %s",
|
||||
e,
|
||||
)
|
||||
_tai_bigram_impl = None
|
||||
return _python_convert_to_bigram_key(tokens)
|
||||
|
||||
87
test/registered/unit/mem_cache/test_convert_to_bigram_key.py
Normal file
87
test/registered/unit/mem_cache/test_convert_to_bigram_key.py
Normal file
@@ -0,0 +1,87 @@
|
||||
"""Unit tests for sglang.srt.mem_cache.utils.convert_to_bigram_key.
|
||||
|
||||
Covers both the pure-Python fallback and the tai-kernel pybind11 path
|
||||
(`tai_kernel.radix.convert_to_bigram_key`) when the extension is importable.
|
||||
|
||||
Run:
|
||||
python -m pytest test/registered/unit/mem_cache/test_convert_to_bigram_key.py -v
|
||||
"""
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
import unittest.mock
|
||||
|
||||
for _mod in ("sgl_kernel", "sgl_kernel.kvcacheio"):
|
||||
if _mod not in sys.modules:
|
||||
sys.modules[_mod] = unittest.mock.MagicMock()
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.mem_cache import utils as mem_cache_utils
|
||||
from sglang.srt.mem_cache.utils import (
|
||||
_python_convert_to_bigram_key,
|
||||
convert_to_bigram_key,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=2, suite="stage-b-test-1-gpu-small")
|
||||
register_amd_ci(est_time=2, suite="stage-b-test-1-gpu-small-amd")
|
||||
|
||||
|
||||
class TestConvertToBigramKey(unittest.TestCase):
|
||||
"""The dispatcher must agree with the pure-Python reference on every input,
|
||||
regardless of whether the tai-kernel extension is loaded."""
|
||||
|
||||
cases = [
|
||||
[],
|
||||
[42],
|
||||
[1, 2],
|
||||
[1, 2, 3, 4, 5],
|
||||
list(range(1024)),
|
||||
# Already-bigram input should be returned as-is (no re-pairing).
|
||||
[(1, 2), (2, 3), (3, 4)],
|
||||
]
|
||||
|
||||
def _assert_equivalent(self, tokens):
|
||||
expected = _python_convert_to_bigram_key(list(tokens))
|
||||
got = convert_to_bigram_key(list(tokens))
|
||||
self.assertEqual(got, expected)
|
||||
self.assertIsInstance(got, list)
|
||||
for item in got:
|
||||
self.assertIsInstance(item, tuple)
|
||||
self.assertEqual(len(item), 2)
|
||||
|
||||
def test_python_fallback(self):
|
||||
# Force the fallback regardless of whether tai_kernel is available.
|
||||
with unittest.mock.patch.object(
|
||||
envs.SGLANG_DISABLE_TAI_BIGRAM, "get", return_value=True
|
||||
):
|
||||
for tokens in self.cases:
|
||||
with self.subTest(tokens=tokens):
|
||||
self._assert_equivalent(tokens)
|
||||
|
||||
def test_tai_kernel_path(self):
|
||||
if mem_cache_utils._tai_bigram_impl is None:
|
||||
self.skipTest("tai_kernel.radix not importable in this environment")
|
||||
with unittest.mock.patch.object(
|
||||
envs.SGLANG_DISABLE_TAI_BIGRAM, "get", return_value=False
|
||||
):
|
||||
for tokens in self.cases:
|
||||
with self.subTest(tokens=tokens):
|
||||
self._assert_equivalent(tokens)
|
||||
|
||||
def test_distinct_tuple_identity(self):
|
||||
# Bigrams that share token values must still be distinct objects so
|
||||
# downstream radix dicts don't alias keys.
|
||||
for force_fallback in (True, False):
|
||||
if not force_fallback and mem_cache_utils._tai_bigram_impl is None:
|
||||
continue
|
||||
with unittest.mock.patch.object(
|
||||
envs.SGLANG_DISABLE_TAI_BIGRAM, "get", return_value=force_fallback
|
||||
):
|
||||
got = convert_to_bigram_key([7, 7, 7])
|
||||
self.assertEqual(got, [(7, 7), (7, 7)])
|
||||
self.assertIsNot(got[0], got[1])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user