tokenizer: Add fastokens support (#23753)

This commit is contained in:
AlonKejzman
2026-04-28 11:43:10 -07:00
committed by GitHub
parent ad785a2299
commit 66ea0aee7f
11 changed files with 152 additions and 4 deletions
@@ -114,6 +114,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>` auto`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>slow</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --tokenizer-backend`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Tokenizer backend. 'huggingface' uses the default HuggingFace tokenizers library; 'fastokens' uses the <a href="https://github.com/crusoecloud/fastokens">fastokens</a> library for faster tokenization. Requires the <code>fastokens</code> package to be installed.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>` huggingface`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>huggingface</code>, <code>fastokens</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --tokenizer-worker-num`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>The worker num of the tokenizer manager.</td>
+5
View File
@@ -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",
]
@@ -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
@@ -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):
+2
View File
@@ -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.
@@ -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
+2
View File
@@ -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
+10
View File
@@ -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,
@@ -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
@@ -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
# ---------------------------------------------------------------------------
@@ -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()