diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py index eef80dc7a..6f1aa28e0 100644 --- a/python/sglang/srt/utils/hf_transformers/common.py +++ b/python/sglang/srt/utils/hf_transformers/common.py @@ -185,6 +185,33 @@ def _resolve_local_or_cached_file(model_name_or_path, filename, revision=None): ) +def _cached_file_exists(model_name_or_path, filename, revision=None) -> bool: + """Whether *filename* is available locally or in the HF cache (no network).""" + try: + _resolve_local_or_cached_file(model_name_or_path, filename, revision) + return True + except Exception: + return False + + +def _remote_file_exists(repo_id, filename, revision=None) -> bool: + """Whether *filename* exists on the HF hub (HEAD request only, no download). + + Returns False on any error (offline, gated, network, invalid id) so callers + fall back to their default path instead of crashing. + """ + from huggingface_hub.constants import HF_HUB_OFFLINE + + if HF_HUB_OFFLINE: + return False + try: + from huggingface_hub import HfApi + + return HfApi().file_exists(repo_id, filename, revision=revision) + except Exception: + return False + + def check_gguf_file(model: Union[str, os.PathLike]) -> bool: model = Path(model) if not model.is_file(): diff --git a/python/sglang/srt/utils/hf_transformers/mistral_utils.py b/python/sglang/srt/utils/hf_transformers/mistral_utils.py index d44fb8669..f75d65dca 100644 --- a/python/sglang/srt/utils/hf_transformers/mistral_utils.py +++ b/python/sglang/srt/utils/hf_transformers/mistral_utils.py @@ -11,7 +11,12 @@ from transformers import AutoConfig, PretrainedConfig, WhisperConfig from sglang.srt.utils import logger -from .common import _ensure_sub_configs, download_from_hf +from .common import ( + _cached_file_exists, + _ensure_sub_configs, + _remote_file_exists, + download_from_hf, +) def adapt_config_dict( @@ -430,6 +435,34 @@ _MISTRAL_TOKENIZER_REDIRECTS = { } +def is_bare_tekken_checkpoint(tokenizer_name, revision=None) -> bool: + """True iff the checkpoint ships tekken.json but no tokenizer.json. + + AutoTokenizer converts tekken.json on the fly, but the converter assigns + BPE ids from rank 0, dropping the 1000 special-token slots that precede + the BPE vocab in tekken's id space — every encoded id is shifted and + generation produces garbage. Such checkpoints must load through the + mistral-common backed tokenizer instead. + """ + + local_dir = Path(tokenizer_name) + if local_dir.is_dir(): + return (local_dir / "tekken.json").is_file() and not ( + local_dir / "tokenizer.json" + ).is_file() + + if _cached_file_exists(tokenizer_name, "tokenizer.json", revision): + return False + if _cached_file_exists(tokenizer_name, "tekken.json", revision): + return True + + # Cold cache: the tokenizer loads before weights, so tekken.json isn't + # cached yet on a first launch — HEAD-probe the hub to still detect it. + if not _remote_file_exists(tokenizer_name, "tekken.json", revision): + return False + return not _remote_file_exists(tokenizer_name, "tokenizer.json", revision) + + def retry_without_mistral_common_kwargs(tokenizer_name, *args, **common_kwargs): """Retry ``AutoTokenizer.from_pretrained`` without kwargs that MistralCommon rejects. diff --git a/python/sglang/srt/utils/hf_transformers/tokenizer.py b/python/sglang/srt/utils/hf_transformers/tokenizer.py index 40b1693a3..def8d2581 100644 --- a/python/sglang/srt/utils/hf_transformers/tokenizer.py +++ b/python/sglang/srt/utils/hf_transformers/tokenizer.py @@ -38,6 +38,7 @@ from .common import ( ) from .mistral_utils import ( _MISTRAL_TOKENIZER_REDIRECTS, + is_bare_tekken_checkpoint, patch_mistral_common_tokenizer, retry_without_mistral_common_kwargs, ) @@ -496,21 +497,38 @@ def get_tokenizer( ) try: - tokenizer = _auto_tokenizer_from_pretrained( - tokenizer_name, *args, **common_kwargs - ) + if is_bare_tekken_checkpoint(tokenizer_name, tokenizer_revision): + from transformers.tokenization_mistral_common import ( + MistralCommonTokenizer, + ) - # 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( + logger.info( + "Detected bare-tekken checkpoint %s (tekken.json, no " + "tokenizer.json); loading via mistral-common MistralCommonTokenizer, " + "ignoring tokenizer_backend=%r.", + tokenizer_name, + tokenizer_backend, + ) + + tokenizer = MistralCommonTokenizer.from_pretrained( + tokenizer_name, revision=tokenizer_revision + ) + else: + tokenizer = _auto_tokenizer_from_pretrained( 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) except Exception as e: if tokenizer_backend == "fastokens": diff --git a/test/registered/unit/tokenizer/test_tekken_tokenizer_routing.py b/test/registered/unit/tokenizer/test_tekken_tokenizer_routing.py new file mode 100644 index 000000000..e96ce850d --- /dev/null +++ b/test/registered/unit/tokenizer/test_tekken_tokenizer_routing.py @@ -0,0 +1,48 @@ +"""Unit tests for bare-tekken checkpoint tokenizer routing — no server, no model loading.""" + +import os +import shutil +import tempfile +import unittest + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=60, suite="base-a-test-cpu") + +from sglang.srt.utils.hf_transformers.mistral_utils import is_bare_tekken_checkpoint +from sglang.srt.utils.hf_transformers.tokenizer import get_tokenizer + +TEKKEN_REPO = "mistralai/Leanstral-1.5-119B-A6B" +PROMPT = "The capital of France is" +# Reference ids from mistral-common, which owns the tekken id space +# (1000 special-token slots precede the BPE vocab). +EXPECTED_IDS = [1784, 8961, 1307, 5498, 1395] + + +class TestBareTekkenDetection(CustomTestCase): + def test_detects_bare_tekken_dir(self): + with tempfile.TemporaryDirectory() as d: + self.assertFalse(is_bare_tekken_checkpoint(d)) + with open(os.path.join(d, "tekken.json"), "w") as f: + f.write("{}") + self.assertTrue(is_bare_tekken_checkpoint(d)) + with open(os.path.join(d, "tokenizer.json"), "w") as f: + f.write("{}") + self.assertFalse(is_bare_tekken_checkpoint(d)) + + +class TestTekkenRouting(CustomTestCase): + def test_get_tokenizer_matches_mistral_common(self): + from huggingface_hub import hf_hub_download + + tekken = hf_hub_download(TEKKEN_REPO, "tekken.json") + with tempfile.TemporaryDirectory() as d: + shutil.copy(tekken, os.path.join(d, "tekken.json")) + tokenizer = get_tokenizer(d) + ids = tokenizer.encode(PROMPT, add_special_tokens=False) + self.assertEqual(ids, EXPECTED_IDS) + + +if __name__ == "__main__": + unittest.main()