Files
2026-08-11 04:33:10 +00:00

80 lines
2.3 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
This test file includes some cases where it is inappropriate to
only get the `eos_token_id` from the tokenizer as defined by
`BaseRenderer.get_eos_token_id`.
"""
from types import SimpleNamespace
from typing import cast
from unittest.mock import patch
from transformers import PretrainedConfig
from vllm.config.model import ModelConfig
from vllm.tokenizers import get_tokenizer
from vllm.transformers_utils import config as config_module
from vllm.transformers_utils.config import try_get_generation_config
def test_get_llama3_eos_token():
model_name = "meta-llama/Llama-3.2-1B-Instruct"
tokenizer = get_tokenizer(model_name)
assert tokenizer.eos_token_id == 128009
generation_config = try_get_generation_config(model_name, trust_remote_code=False)
assert generation_config is not None
assert generation_config.eos_token_id == [128001, 128008, 128009]
def test_get_blip2_eos_token():
model_name = "Salesforce/blip2-opt-2.7b"
tokenizer = get_tokenizer(model_name)
assert tokenizer.eos_token_id == 2
generation_config = try_get_generation_config(model_name, trust_remote_code=False)
assert generation_config is not None
assert generation_config.eos_token_id == 50118
def test_model_config_generation_fallback_forwards_code_revision():
model_config = cast(
ModelConfig,
SimpleNamespace(
generation_config="auto",
hf_config_path=None,
model="org/model",
trust_remote_code=True,
revision="model-pin",
code_revision="code-pin",
config_format="auto",
hf_token=None,
),
)
with (
patch.object(
config_module.GenerationConfig,
"from_pretrained",
side_effect=OSError,
),
patch.object(
config_module,
"get_config",
return_value=PretrainedConfig(),
) as get_config,
):
ModelConfig.try_get_generation_config(model_config)
get_config.assert_called_once_with(
"org/model",
trust_remote_code=True,
revision="model-pin",
code_revision="code-pin",
config_format="auto",
token=None,
)