forked from Karylab-cklius/vllm
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:
co-authored by
OpenAI Codex
parent
9d808e2309
commit
a9a8a32dcd
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user