diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx index 71a7c1ac9..4285ff521 100644 --- a/docs_new/docs/advanced_features/server_arguments.mdx +++ b/docs_new/docs/advanced_features/server_arguments.mdx @@ -114,6 +114,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s ` auto` auto, slow + + ` --tokenizer-backend` + Tokenizer backend. 'huggingface' uses the default HuggingFace tokenizers library; 'fastokens' uses the fastokens library for faster tokenization. Requires the fastokens package to be installed. + ` huggingface` + huggingface, fastokens + ` --tokenizer-worker-num` The worker num of the tokenizer manager. diff --git a/python/pyproject.toml b/python/pyproject.toml index 2b81b52b4..a906ebf9d 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -142,6 +142,10 @@ http2 = [ "granian>=2.6.0", ] +fastokens = [ + "fastokens>=0.1.1,<0.2.0", +] + test = [ "accelerate", "addict", @@ -158,6 +162,7 @@ test = [ "pytest-cov", "diff-cover", "sentence_transformers", + "sglang[fastokens]", "tabulate", "granian>=2.6.0", ] diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index 40dd04adc..6047b434b 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -662,6 +662,7 @@ class MMReceiverBase(ABC): trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, use_fast=not server_args.disable_fast_image_processor, + tokenizer_backend=server_args.tokenizer_backend, ) except ValueError as e: error_message = str(e) @@ -675,6 +676,7 @@ class MMReceiverBase(ABC): trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, use_fast=True, + tokenizer_backend=server_args.tokenizer_backend, ) else: raise e diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 9749af0e0..b760c2a6a 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -108,6 +108,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, + tokenizer_backend=server_args.tokenizer_backend, ) def init_running_status(self, server_args: ServerArgs): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 47e8f659b..225b0ee66 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -556,6 +556,7 @@ class Scheduler( trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, use_fast=not server_args.disable_fast_image_processor, + tokenizer_backend=server_args.tokenizer_backend, ) self.tokenizer = get_tokenizer_from_processor(self.processor) else: @@ -564,6 +565,7 @@ class Scheduler( tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, + tokenizer_backend=server_args.tokenizer_backend, ) # Load multimodal processor for M-RoPE fallback computation. diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 9e75b71d7..66f379a69 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -324,6 +324,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, + tokenizer_backend=server_args.tokenizer_backend, ) # Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal) @@ -2691,6 +2692,7 @@ def _get_processor_wrapper(server_args): trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, use_fast=not server_args.disable_fast_image_processor, + tokenizer_backend=server_args.tokenizer_backend, ) except ValueError as e: error_message = str(e) @@ -2704,6 +2706,7 @@ def _get_processor_wrapper(server_args): trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, use_fast=True, + tokenizer_backend=server_args.tokenizer_backend, ) else: raise e diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index c83053da9..efb7f7563 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -273,6 +273,7 @@ class TpModelWorker(BaseTpWorker): tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, + tokenizer_backend=server_args.tokenizer_backend, ) self.tokenizer = get_tokenizer_from_processor(self.processor) else: @@ -281,6 +282,7 @@ class TpModelWorker(BaseTpWorker): tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, + tokenizer_backend=server_args.tokenizer_backend, ) self.device = self.model_runner.device diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 35a07df22..817ae43fb 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -304,6 +304,7 @@ class ServerArgs: model_path: str tokenizer_path: Optional[str] = None tokenizer_mode: str = "auto" + tokenizer_backend: str = "huggingface" tokenizer_worker_num: int = 1 skip_tokenizer_init: bool = False load_format: str = "auto" @@ -4103,6 +4104,15 @@ class ServerArgs: "tokenizer if available, and 'slow' will " "always use the slow tokenizer.", ) + parser.add_argument( + "--tokenizer-backend", + type=str, + default=ServerArgs.tokenizer_backend, + choices=["huggingface", "fastokens"], + help="Tokenizer backend. 'huggingface' uses the default HuggingFace " + "tokenizers library, and 'fastokens' uses the fastokens library " + "for faster tokenization. Requires the fastokens package to be installed.", + ) parser.add_argument( "--tokenizer-worker-num", type=int, diff --git a/python/sglang/srt/utils/hf_transformers/processor.py b/python/sglang/srt/utils/hf_transformers/processor.py index 31d5905a8..e19227a3e 100644 --- a/python/sglang/srt/utils/hf_transformers/processor.py +++ b/python/sglang/srt/utils/hf_transformers/processor.py @@ -141,8 +141,14 @@ def get_processor( trust_remote_code: bool = False, tokenizer_revision: Optional[str] = None, use_fast: Optional[bool] = True, + tokenizer_backend: str = "huggingface", **kwargs, ): + if tokenizer_backend == "fastokens": + from .tokenizer import _ensure_fastokens_patched + + _ensure_fastokens_patched() + revision = kwargs.pop("revision", tokenizer_revision) if is_mistral_model(tokenizer_name): config = load_mistral_config( @@ -266,6 +272,7 @@ def get_processor( tokenizer_mode=tokenizer_mode, trust_remote_code=trust_remote_code, tokenizer_revision=revision, + tokenizer_backend=tokenizer_backend, ) if isinstance(processor, PreTrainedTokenizerBase): processor = tokenizer diff --git a/python/sglang/srt/utils/hf_transformers/tokenizer.py b/python/sglang/srt/utils/hf_transformers/tokenizer.py index b965b804e..41df30610 100644 --- a/python/sglang/srt/utils/hf_transformers/tokenizer.py +++ b/python/sglang/srt/utils/hf_transformers/tokenizer.py @@ -436,20 +436,46 @@ def _apply_post_load_fixes(tokenizer, tokenizer_name, revision): # --------------------------------------------------------------------------- +_fastokens_patched = False + + +def _ensure_fastokens_patched(): + """Monkey-patch transformers to use the fastokens backend (once).""" + global _fastokens_patched + if _fastokens_patched: + return + try: + import fastokens + except ImportError: + raise ImportError( + "The fastokens package is required when --tokenizer-backend=fastokens. " + "Install it with: pip install 'sglang[fastokens]'" + ) from None + + fastokens.patch_transformers() + _fastokens_patched = True + logger.info("fastokens backend enabled - transformers patched successfully") + + def get_tokenizer( tokenizer_name: str, *args, tokenizer_mode: str = "auto", trust_remote_code: bool = False, tokenizer_revision: Optional[str] = None, + tokenizer_backend: str = "huggingface", **kwargs, ) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast]: """Gets a tokenizer for the given model name via Huggingface.""" + # Tiktoken format has its own backend — no fastokens patching needed. if tokenizer_name.endswith(".json"): from sglang.srt.tokenizer.tiktoken_tokenizer import TiktokenTokenizer return TiktokenTokenizer(tokenizer_name) + if tokenizer_backend == "fastokens": + _ensure_fastokens_patched() + if tokenizer_mode == "slow": if kwargs.get("use_fast", False): raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.") @@ -470,12 +496,32 @@ def get_tokenizer( **kwargs, ) - tokenizer = _auto_tokenizer_from_pretrained(tokenizer_name, *args, **common_kwargs) + try: + tokenizer = _auto_tokenizer_from_pretrained( + tokenizer_name, *args, **common_kwargs + ) - if type(tokenizer).__name__ == _TOKENIZERS_BACKEND: - tokenizer = _resolve_tokenizers_backend(tokenizer_name, *args, **common_kwargs) + # With fastokens, the patched TokenizersBackend.from_pretrained already + # returned a tokenizer whose backend is a fastokens shim. Re-resolving via + # the declared class (e.g. Qwen2Tokenizer) would discard that work. + if ( + type(tokenizer).__name__ == _TOKENIZERS_BACKEND + and tokenizer_backend != "fastokens" + ): + tokenizer = _resolve_tokenizers_backend( + tokenizer_name, *args, **common_kwargs + ) - return _apply_post_load_fixes(tokenizer, tokenizer_name, tokenizer_revision) + return _apply_post_load_fixes(tokenizer, tokenizer_name, tokenizer_revision) + except Exception as e: + if tokenizer_backend == "fastokens": + raise RuntimeError( + f"fastokens failed to load tokenizer for {tokenizer_name!r}. " + f"This model's tokenizer may not be supported by fastokens — " + f"see https://github.com/crusoecloud/fastokens. " + f"Re-run without --tokenizer-backend=fastokens to use the default backend." + ) from e + raise # --------------------------------------------------------------------------- diff --git a/test/registered/unit/utils/test_hf_transformers_fastokens.py b/test/registered/unit/utils/test_hf_transformers_fastokens.py new file mode 100644 index 000000000..9d63e91d5 --- /dev/null +++ b/test/registered/unit/utils/test_hf_transformers_fastokens.py @@ -0,0 +1,64 @@ +"""End-to-end verification that --tokenizer-backend=fastokens swaps the +backend of the loaded tokenizer with fastokens' _TokenizerShim. +""" + +import unittest + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN, + CustomTestCase, +) + +TOKENIZER_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN + +register_cpu_ci(est_time=30, suite="stage-a-test-cpu") + + +try: + import fastokens # noqa: F401 + + HAS_FASTOKENS = True +except ImportError: + HAS_FASTOKENS = False + + +@unittest.skipUnless(HAS_FASTOKENS, "fastokens package not installed") +class TestFastokensBackend(CustomTestCase): + def test_shim_is_applied(self): + # `_TokenizerShim` is fastokens' private compat shim. SGLang's + # integration relies on `tokenizer._tokenizer` being an instance of + # this class to confirm fastokens is wired up. If fastokens renames + # or restructures it, update both this assertion and any code in + # SGLang that depends on the same private name. + from fastokens._compat import _TokenizerShim + + from sglang.srt.utils.hf_transformers.tokenizer import get_tokenizer + + tokenizer = get_tokenizer( + TOKENIZER_MODEL, + tokenizer_backend="fastokens", + ) + backend = getattr(tokenizer, "_tokenizer", None) + self.assertIsInstance( + backend, + _TokenizerShim, + f"Expected tokenizer._tokenizer to be _TokenizerShim, " + f"got {type(backend).__name__}", + ) + + def test_encode_decode_roundtrip(self): + from sglang.srt.utils.hf_transformers.tokenizer import get_tokenizer + + tokenizer = get_tokenizer( + TOKENIZER_MODEL, + tokenizer_backend="fastokens", + ) + text = "Hello, world!" + ids = tokenizer.encode(text, add_special_tokens=False) + self.assertGreater(len(ids), 0) + self.assertEqual(tokenizer.decode(ids, skip_special_tokens=True), text) + + +if __name__ == "__main__": + unittest.main()