tokenizer: Add fastokens support (#23753)
This commit is contained in:
@@ -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>
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user