Register parsed config classes before tokenizer init (#40299)

Signed-off-by: Bortlesboat <bortstheboat@gmail.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
This commit is contained in:
Andrew Barnes
2026-06-16 05:33:08 +00:00
committed by GitHub
co-authored by OpenAI Codex
parent 9d808e2309
commit a9a8a32dcd
3 changed files with 85 additions and 5 deletions
+64
View File
@@ -1,15 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import json
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
import pytest
from transformers import AutoConfig
from transformers.models.auto.configuration_auto import CONFIG_MAPPING
from vllm.tokenizers import TokenizerLike
from vllm.tokenizers.registry import (
TokenizerRegistry,
cached_get_tokenizer,
cached_resolve_tokenizer_args,
cached_tokenizer_from_config,
get_tokenizer,
resolve_tokenizer_args,
)
from vllm.transformers_utils.configs.qwen3_5_moe import Qwen3_5MoeConfig
class TestTokenizer(TokenizerLike):
@@ -75,3 +84,58 @@ def test_customized_tokenizer():
assert tokenizer.bos_token_id == 0
assert tokenizer.eos_token_id == 1
assert tokenizer.pad_token_id == 2
def test_cached_tokenizer_from_config_registers_local_config(tmp_path: Path):
(tmp_path / "config.json").write_text(
json.dumps({"model_type": "qwen3_5_moe"}),
encoding="utf-8",
)
model_config = SimpleNamespace(
skip_tokenizer_init=False,
tokenizer=str(tmp_path),
runner_type="generate",
tokenizer_mode="hf",
tokenizer_revision=None,
trust_remote_code=True,
hf_config=Qwen3_5MoeConfig(),
)
registered_config = CONFIG_MAPPING._extra_content.pop("qwen3_5_moe", None)
cached_get_tokenizer.cache_clear()
cached_resolve_tokenizer_args.cache_clear()
try:
def fake_from_pretrained(path_or_repo_id: str, *args, **kwargs):
loaded_config = AutoConfig.from_pretrained(
path_or_repo_id,
trust_remote_code=False,
)
assert isinstance(loaded_config, Qwen3_5MoeConfig)
return SimpleNamespace(is_fast=True)
with (
patch(
"vllm.tokenizers.registry.logger.debug_once",
lambda *args, **kwargs: None,
),
patch(
"vllm.tokenizers.hf.AutoTokenizer.from_pretrained",
side_effect=fake_from_pretrained,
),
patch(
"vllm.tokenizers.hf.get_cached_tokenizer",
side_effect=lambda tokenizer: tokenizer,
),
):
tokenizer = cached_tokenizer_from_config(model_config)
assert tokenizer.is_fast is True
finally:
cached_get_tokenizer.cache_clear()
cached_resolve_tokenizer_args.cache_clear()
CONFIG_MAPPING._extra_content.pop("qwen3_5_moe", None)
if registered_config is not None:
CONFIG_MAPPING._extra_content["qwen3_5_moe"] = registered_config
+3 -1
View File
@@ -11,7 +11,7 @@ from typing_extensions import TypeVar, assert_never
import vllm.envs as envs
from vllm.logger import init_logger
from vllm.transformers_utils.config import get_config
from vllm.transformers_utils.config import _maybe_register_hf_config, get_config
from vllm.transformers_utils.repo_utils import (
any_pattern_in_repo_files,
is_mistral_model_repo,
@@ -246,6 +246,8 @@ def cached_tokenizer_from_config(model_config: "ModelConfig", **kwargs):
if model_config.skip_tokenizer_init:
return None
_maybe_register_hf_config(getattr(model_config, "hf_config", None))
return cached_get_tokenizer(
model_config.tokenizer,
runner_type=model_config.runner_type,
+18 -4
View File
@@ -141,6 +141,22 @@ _AUTO_CONFIG_KWARGS_OVERRIDES: dict[str, dict[str, Any]] = {
}
def _register_config_class(
model_type: str, config_class: type[PretrainedConfig]
) -> None:
config_class.model_type = model_type
AutoConfig.register(model_type, config_class, exist_ok=True)
def _maybe_register_hf_config(config: PretrainedConfig | None) -> None:
if config is None:
return
model_type = getattr(config, "model_type", None)
if isinstance(model_type, str) and model_type in _CONFIG_REGISTRY:
_register_config_class(model_type, _CONFIG_REGISTRY[model_type])
def is_rope_parameters_nested(rope_parameters: dict[str, Any]) -> bool:
"""Check if rope_parameters is nested by layer types."""
# Cannot be nested if rope_parameters is empty
@@ -244,8 +260,7 @@ class HFConfigParser(ConfigParserBase):
# in future calls to `from_pretrained` (e.g. from
# AutoTokenizer or AutoProcessor).
config_class = _CONFIG_REGISTRY[model_type]
config_class.model_type = model_type
AutoConfig.register(model_type, config_class, exist_ok=True)
_register_config_class(model_type, config_class)
# If the on-disk model_type differs from the overridden
# one, register under both so AutoConfig.from_pretrained
# returns the correct class regardless of what the
@@ -253,8 +268,7 @@ class HFConfigParser(ConfigParserBase):
if (
config_model_type := config_dict.get("model_type")
) and config_model_type != model_type:
config_class.model_type = config_model_type
AutoConfig.register(config_model_type, config_class, exist_ok=True)
_register_config_class(config_model_type, config_class)
config_class.model_type = model_type
# Now that it is registered, it is not considered remote code anymore
trust_remote_code = False