From 12df16607b3d74f21caf62d68a9882abe3f1e901 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 5 Jan 2026 09:12:05 +0800 Subject: [PATCH] Tiny speed up kimi detokenizer by 10x (#16427) --- python/sglang/srt/environ.py | 3 + .../sglang/srt/utils/hf_transformers_utils.py | 2 + python/sglang/srt/utils/patch_tokenizer.py | 116 ++++++++++++ test/srt/test_patch_tokenizer.py | 175 ++++++++++++++++++ 4 files changed, 296 insertions(+) create mode 100644 python/sglang/srt/utils/patch_tokenizer.py create mode 100644 test/srt/test_patch_tokenizer.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index a5f1d21e4..de90328ef 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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 diff --git a/python/sglang/srt/utils/hf_transformers_utils.py b/python/sglang/srt/utils/hf_transformers_utils.py index 04ea73141..f88c4889d 100644 --- a/python/sglang/srt/utils/hf_transformers_utils.py +++ b/python/sglang/srt/utils/hf_transformers_utils.py @@ -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 diff --git a/python/sglang/srt/utils/patch_tokenizer.py b/python/sglang/srt/utils/patch_tokenizer.py new file mode 100644 index 000000000..7ad1114da --- /dev/null +++ b/python/sglang/srt/utils/patch_tokenizer.py @@ -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 diff --git a/test/srt/test_patch_tokenizer.py b/test/srt/test_patch_tokenizer.py new file mode 100644 index 000000000..cc34bc040 --- /dev/null +++ b/test/srt/test_patch_tokenizer.py @@ -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": ""}) + 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([""], special_tokens=True) + self.assertIn("Cannot add special tokens after patch", str(ctx.exception)) + + tokenizer.add_tokens([""], 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()