Tiny speed up kimi detokenizer by 10x (#16427)
This commit is contained in:
@@ -416,6 +416,9 @@ class Envs:
|
||||
SGLANG_ENABLE_METRICS_DEVICE_TIMER = EnvBool(False)
|
||||
SGLANG_ENABLE_METRICS_DP_ATTENTION = EnvBool(False)
|
||||
|
||||
# Tokenizer
|
||||
SGLANG_PATCH_TOKENIZER = EnvBool(False) # TODO enable by default
|
||||
|
||||
# fmt: on
|
||||
|
||||
|
||||
|
||||
@@ -68,6 +68,7 @@ from sglang.srt.configs.internvl import InternVLChatConfig
|
||||
from sglang.srt.connector import create_remote_connector
|
||||
from sglang.srt.multimodal.customized_mm_processor_utils import _CUSTOMIZED_MM_PROCESSOR
|
||||
from sglang.srt.utils import is_remote_url, logger, lru_cache_frozenset, mistral_utils
|
||||
from sglang.srt.utils.patch_tokenizer import patch_tokenizer
|
||||
|
||||
_CONFIG_REGISTRY: List[Type[PretrainedConfig]] = [
|
||||
ChatGLMConfig,
|
||||
@@ -501,6 +502,7 @@ def get_tokenizer(
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
tokenizer = patch_tokenizer(tokenizer)
|
||||
return tokenizer
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
import logging
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def patch_tokenizer(tokenizer):
|
||||
if not envs.SGLANG_PATCH_TOKENIZER.get():
|
||||
return tokenizer
|
||||
|
||||
if _is_kimi_tiktoken_tokenizer(tokenizer):
|
||||
logger.info(
|
||||
f"Applying special tokens cache patch for Kimi tokenizer: {type(tokenizer)}"
|
||||
)
|
||||
return _SpecialTokensCachePatcher.patch(tokenizer)
|
||||
|
||||
return tokenizer
|
||||
|
||||
|
||||
def unpatch_tokenizer(tokenizer):
|
||||
return _SpecialTokensCachePatcher.unpatch(tokenizer)
|
||||
|
||||
|
||||
def _is_kimi_tiktoken_tokenizer(tokenizer):
|
||||
cls = type(tokenizer)
|
||||
class_name = cls.__name__
|
||||
module_name = cls.__module__ or ""
|
||||
return class_name == "TikTokenTokenizer" and "tokenization_kimi" in module_name
|
||||
|
||||
|
||||
class _SpecialTokensCachePatcher:
|
||||
_PATCHED_FLAG = "_sglang_special_tokens_patched"
|
||||
_CACHED_TOKENS_ATTR = "_sglang_cached_special_tokens"
|
||||
_CACHED_IDS_ATTR = "_sglang_cached_special_ids"
|
||||
|
||||
@classmethod
|
||||
def patch(cls, tokenizer):
|
||||
tokenizer_cls = type(tokenizer)
|
||||
|
||||
if getattr(tokenizer_cls, cls._PATCHED_FLAG, False):
|
||||
return tokenizer
|
||||
|
||||
tokenizer_cls._original_all_special_tokens = (
|
||||
tokenizer_cls.all_special_tokens.fget
|
||||
)
|
||||
tokenizer_cls._original_all_special_ids = tokenizer_cls.all_special_ids.fget
|
||||
tokenizer_cls._original_add_special_tokens = tokenizer_cls.add_special_tokens
|
||||
tokenizer_cls._original_add_tokens = tokenizer_cls.add_tokens
|
||||
|
||||
patched_all_special_tokens = _make_cached_property(
|
||||
cls._CACHED_TOKENS_ATTR, tokenizer_cls._original_all_special_tokens
|
||||
)
|
||||
patched_all_special_ids = _make_cached_property(
|
||||
cls._CACHED_IDS_ATTR, tokenizer_cls._original_all_special_ids
|
||||
)
|
||||
|
||||
def patched_add_special_tokens(self, *args, **kwargs):
|
||||
assert (
|
||||
False
|
||||
), "Cannot modify special tokens after patch. Call unpatch_tokenizer first."
|
||||
|
||||
def patched_add_tokens(self, new_tokens, special_tokens=False):
|
||||
assert (
|
||||
not special_tokens
|
||||
), "Cannot add special tokens after patch. Call unpatch_tokenizer first."
|
||||
return tokenizer_cls._original_add_tokens(
|
||||
self, new_tokens, special_tokens=False
|
||||
)
|
||||
|
||||
tokenizer_cls.all_special_tokens = patched_all_special_tokens
|
||||
tokenizer_cls.all_special_ids = patched_all_special_ids
|
||||
tokenizer_cls.add_special_tokens = patched_add_special_tokens
|
||||
tokenizer_cls.add_tokens = patched_add_tokens
|
||||
setattr(tokenizer_cls, cls._PATCHED_FLAG, True)
|
||||
|
||||
return tokenizer
|
||||
|
||||
@classmethod
|
||||
def unpatch(cls, tokenizer):
|
||||
tokenizer_cls = type(tokenizer)
|
||||
|
||||
if not getattr(tokenizer_cls, cls._PATCHED_FLAG, False):
|
||||
return tokenizer
|
||||
|
||||
tokenizer_cls.all_special_tokens = property(
|
||||
tokenizer_cls._original_all_special_tokens
|
||||
)
|
||||
tokenizer_cls.all_special_ids = property(
|
||||
tokenizer_cls._original_all_special_ids
|
||||
)
|
||||
tokenizer_cls.add_special_tokens = tokenizer_cls._original_add_special_tokens
|
||||
tokenizer_cls.add_tokens = tokenizer_cls._original_add_tokens
|
||||
|
||||
del tokenizer_cls._original_all_special_tokens
|
||||
del tokenizer_cls._original_all_special_ids
|
||||
del tokenizer_cls._original_add_special_tokens
|
||||
del tokenizer_cls._original_add_tokens
|
||||
delattr(tokenizer_cls, cls._PATCHED_FLAG)
|
||||
|
||||
for attr in [cls._CACHED_TOKENS_ATTR, cls._CACHED_IDS_ATTR]:
|
||||
if hasattr(tokenizer, attr):
|
||||
delattr(tokenizer, attr)
|
||||
|
||||
logger.info(f"Unpatched special tokens cache for {tokenizer_cls.__name__}")
|
||||
return tokenizer
|
||||
|
||||
|
||||
def _make_cached_property(cache_attr, original_fn):
|
||||
@property
|
||||
def cached_prop(self):
|
||||
if getattr(self, cache_attr, None) is None:
|
||||
setattr(self, cache_attr, original_fn(self))
|
||||
return getattr(self, cache_attr)
|
||||
|
||||
return cached_prop
|
||||
@@ -0,0 +1,175 @@
|
||||
import random
|
||||
import unittest
|
||||
from contextlib import contextmanager
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from sglang.srt.utils.patch_tokenizer import (
|
||||
_SpecialTokensCachePatcher,
|
||||
unpatch_tokenizer,
|
||||
)
|
||||
|
||||
|
||||
class TestPatchTokenizerEndToEndTest(unittest.TestCase):
|
||||
def test_patched_produces_same_results_as_raw(self):
|
||||
tokenizer = _load_tokenizer()
|
||||
test_texts = self._generate_test_texts(tokenizer)
|
||||
raw_results = self._run_tokenizer_ops(tokenizer, test_texts)
|
||||
|
||||
_SpecialTokensCachePatcher.patch(tokenizer)
|
||||
patched_results = self._run_tokenizer_ops(tokenizer, test_texts)
|
||||
unpatch_tokenizer(tokenizer)
|
||||
|
||||
self.assertEqual(raw_results, patched_results)
|
||||
|
||||
@classmethod
|
||||
def _generate_test_texts(cls, tokenizer):
|
||||
special_tokens = tokenizer.all_special_tokens
|
||||
return [
|
||||
"Hello, world!",
|
||||
"This is a longer sentence with multiple words.",
|
||||
"Numbers 12345 and symbols !@#$%",
|
||||
" leading and trailing spaces ",
|
||||
"\n\nMultiple\n\nNewlines\n\n",
|
||||
*[f"Text with {tok} inside" for tok in special_tokens],
|
||||
" ".join(special_tokens),
|
||||
*[
|
||||
cls._random_text_from_tokens(tokenizer, num_tokens=100)
|
||||
for _ in range(5)
|
||||
],
|
||||
*[
|
||||
cls._random_text_from_tokens(tokenizer, num_tokens=1000)
|
||||
for _ in range(3)
|
||||
],
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def _random_text_from_tokens(cls, tokenizer, num_tokens):
|
||||
token_ids = [
|
||||
random.randint(0, tokenizer.vocab_size - 1) for _ in range(num_tokens)
|
||||
]
|
||||
return tokenizer.decode(token_ids)
|
||||
|
||||
@classmethod
|
||||
def _run_tokenizer_ops(cls, tokenizer, texts):
|
||||
encode_results = [tokenizer.encode(t) for t in texts]
|
||||
batch_encode_results = tokenizer(texts)["input_ids"]
|
||||
return {
|
||||
"encode": encode_results,
|
||||
"batch_encode": batch_encode_results,
|
||||
"decode": [
|
||||
tokenizer.decode(ids, skip_special_tokens=True)
|
||||
for ids in encode_results
|
||||
],
|
||||
"batch_decode": tokenizer.batch_decode(
|
||||
encode_results, skip_special_tokens=True
|
||||
),
|
||||
"special_tokens": tokenizer.all_special_tokens,
|
||||
"special_ids": tokenizer.all_special_ids,
|
||||
}
|
||||
|
||||
|
||||
class TestPatchTokenizerUnitTest(unittest.TestCase):
|
||||
def test_patch_unpatch_restores_original(self):
|
||||
tokenizer = _load_tokenizer()
|
||||
cls = type(tokenizer)
|
||||
|
||||
original_ids = _get_class_attr_ids(cls)
|
||||
|
||||
_SpecialTokensCachePatcher.patch(tokenizer)
|
||||
self.assertTrue(getattr(cls, "_sglang_special_tokens_patched", False))
|
||||
|
||||
patched_ids = _get_class_attr_ids(cls)
|
||||
changed_attrs = [
|
||||
name
|
||||
for name in original_ids
|
||||
if name in patched_ids and patched_ids[name] != original_ids[name]
|
||||
]
|
||||
self.assertGreater(len(changed_attrs), 0, "Patch should change some attributes")
|
||||
|
||||
unpatch_tokenizer(tokenizer)
|
||||
self.assertFalse(getattr(cls, "_sglang_special_tokens_patched", False))
|
||||
|
||||
restored_ids = _get_class_attr_ids(cls)
|
||||
for name in original_ids:
|
||||
if name.startswith("_sglang") or name.startswith("_original"):
|
||||
continue
|
||||
self.assertEqual(
|
||||
restored_ids.get(name),
|
||||
original_ids[name],
|
||||
f"Attribute {name} should be restored to original",
|
||||
)
|
||||
|
||||
def test_patch_caches_special_tokens(self):
|
||||
with _patched_tokenizer() as tokenizer:
|
||||
tokens1 = tokenizer.all_special_tokens
|
||||
ids1 = tokenizer.all_special_ids
|
||||
tokens2 = tokenizer.all_special_tokens
|
||||
ids2 = tokenizer.all_special_ids
|
||||
|
||||
self.assertIs(tokens1, tokens2)
|
||||
self.assertIs(ids1, ids2)
|
||||
|
||||
def test_patch_blocks_add_special_tokens(self):
|
||||
with _patched_tokenizer() as tokenizer:
|
||||
with self.assertRaises(AssertionError) as ctx:
|
||||
tokenizer.add_special_tokens({"pad_token": "<pad>"})
|
||||
self.assertIn(
|
||||
"Cannot modify special tokens after patch", str(ctx.exception)
|
||||
)
|
||||
|
||||
def test_patch_blocks_add_tokens_with_special_flag(self):
|
||||
with _patched_tokenizer() as tokenizer:
|
||||
with self.assertRaises(AssertionError) as ctx:
|
||||
tokenizer.add_tokens(["<new>"], special_tokens=True)
|
||||
self.assertIn("Cannot add special tokens after patch", str(ctx.exception))
|
||||
|
||||
tokenizer.add_tokens(["<regular>"], special_tokens=False)
|
||||
|
||||
def test_unpatch_clears_cache(self):
|
||||
with _patched_tokenizer() as tokenizer:
|
||||
_ = tokenizer.all_special_tokens
|
||||
_ = tokenizer.all_special_ids
|
||||
self.assertTrue(hasattr(tokenizer, "_sglang_cached_special_tokens"))
|
||||
self.assertTrue(hasattr(tokenizer, "_sglang_cached_special_ids"))
|
||||
|
||||
self.assertFalse(hasattr(tokenizer, "_sglang_cached_special_tokens"))
|
||||
self.assertFalse(hasattr(tokenizer, "_sglang_cached_special_ids"))
|
||||
|
||||
def test_double_patch_is_idempotent(self):
|
||||
tokenizer = _load_tokenizer()
|
||||
_SpecialTokensCachePatcher.patch(tokenizer)
|
||||
_SpecialTokensCachePatcher.patch(tokenizer)
|
||||
|
||||
self.assertTrue(
|
||||
getattr(type(tokenizer), "_sglang_special_tokens_patched", False)
|
||||
)
|
||||
|
||||
unpatch_tokenizer(tokenizer)
|
||||
|
||||
|
||||
def _get_class_attr_ids(cls):
|
||||
return {
|
||||
n: id(v.fget if isinstance(v, property) else v) for n, v in vars(cls).items()
|
||||
}
|
||||
|
||||
|
||||
def _load_tokenizer():
|
||||
# The slowness is mainly observed in Kimi
|
||||
return AutoTokenizer.from_pretrained(
|
||||
"nvidia/Kimi-K2-Thinking-NVFP4", trust_remote_code=True
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _patched_tokenizer():
|
||||
tokenizer = _load_tokenizer()
|
||||
_SpecialTokensCachePatcher.patch(tokenizer)
|
||||
try:
|
||||
yield tokenizer
|
||||
finally:
|
||||
unpatch_tokenizer(tokenizer)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user