Skip to content

Commit 3f5051b

Browse files
authored
[Fix] Fix vLLM chat template BOS handling (#2554)
1 parent 1e0b930 commit 3f5051b

2 files changed

Lines changed: 47 additions & 7 deletions

File tree

opencompass/models/vllm_with_tf_above_v4_33.py

Lines changed: 4 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -125,13 +125,10 @@ def generate(self, inputs: List[str], max_out_len: int, stopping_criteria: List[
125125
messages = _format_with_fast_chat_template(messages, self.fastchat_template)
126126
else:
127127
messages = [self.tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=False, **self.chat_template_kwargs) for m in messages]
128-
# vLLM tokenize prompts by AutoTokenizer with its default parameter "add_special_token=True"
129-
# OC add bos_token in the prompt, which requires tokenizing prompts using "add_speicial_token=False"
130-
# But vLLM doesn't have "add_speicial_token" in the pipeline API. So, we remove bos_token
131-
# from messages as a workaround
132-
if self.tokenizer.bos_token:
133-
bos_token = self.tokenizer.bos_token
134-
messages = [message.removeprefix(bos_token) if message.startswith(bos_token) else message for message in messages]
128+
messages = [{
129+
'prompt': message,
130+
'prompt_token_ids': self.tokenizer.encode(message, add_special_tokens=False),
131+
} for message in messages]
135132
DEFAULT_GENERATION_KWARGS = {
136133
'temperature': 0,
137134
'max_tokens': max_out_len,

tests/models/test_vllm_with_tf_above_v4_33.py

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,49 @@ def test_generate_basic(self, mock_ray, mock_sampling_params_class,
118118
self.assertEqual(results[0], 'Generated response')
119119
mock_model.generate.assert_called_once()
120120

121+
@patch('opencompass.models.vllm_with_tf_above_v4_33._convert_chat_messages'
122+
)
123+
@patch('opencompass.models.vllm_with_tf_above_v4_33.SamplingParams')
124+
def test_generate_preserves_bos_token_with_token_ids(
125+
self, mock_sampling_params_class, mock_convert_messages):
126+
"""Test BOS token is preserved in vLLM prompts."""
127+
model = VLLMwithChatTemplate.__new__(VLLMwithChatTemplate)
128+
model.fastchat_template = None
129+
model.chat_template_kwargs = {}
130+
model.generation_kwargs = {}
131+
model.stop_words = []
132+
model.lora_path = None
133+
model.logger = MagicMock()
134+
mock_model = MagicMock()
135+
mock_tokenizer = MagicMock()
136+
mock_tokenizer.bos_token = '<bos>'
137+
mock_tokenizer.eos_token = None
138+
mock_tokenizer.apply_chat_template.return_value = '<bos>Formatted prompt'
139+
mock_tokenizer.encode.return_value = [1, 2, 3]
140+
model.tokenizer = mock_tokenizer
141+
model.model = mock_model
142+
143+
mock_convert_messages.return_value = [{
144+
'role': 'user',
145+
'content': 'Hello'
146+
}]
147+
mock_output = MagicMock()
148+
mock_output.outputs = [MagicMock(text='Generated response')]
149+
mock_model.generate.return_value = [mock_output]
150+
mock_sampling_params = MagicMock()
151+
mock_sampling_params_class.return_value = mock_sampling_params
152+
153+
results = model.generate(['Hello'], max_out_len=100)
154+
155+
self.assertEqual(results, ['Generated response'])
156+
mock_tokenizer.encode.assert_called_once_with(
157+
'<bos>Formatted prompt', add_special_tokens=False)
158+
prompts = mock_model.generate.call_args[0][0]
159+
self.assertEqual(prompts, [{
160+
'prompt': '<bos>Formatted prompt',
161+
'prompt_token_ids': [1, 2, 3],
162+
}])
163+
121164
@patch('opencompass.models.vllm_with_tf_above_v4_33.LLM')
122165
@patch('transformers.AutoTokenizer')
123166
@patch(

0 commit comments

Comments
 (0)