Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
201 changes: 201 additions & 0 deletions tests/entrypoints/pooling/embed/test_io_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,17 +3,108 @@
"""Unit tests for EmbedIOProcessor."""

import pytest
from pydantic import TypeAdapter

from vllm import PoolingParams
from vllm.entrypoints.pooling.embed.io_processor import EmbedIOProcessor
from vllm.entrypoints.pooling.embed.protocol import (
CohereEmbedContent,
CohereEmbedInput,
CohereEmbedRequest,
EmbeddingBatchChatInputRequest,
EmbeddingBatchChatRequest,
EmbeddingChatInputRequest,
EmbeddingChatRequest,
EmbeddingCompletionRequest,
EmbeddingRequest,
)
from vllm.entrypoints.pooling.typing import PoolingServeContext


class TestEmbeddingRequestParsing:
"""Unit tests for OpenAI embedding request parsing."""

def test_input_messages_parses_as_chat_request(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"input": [{"role": "user", "content": "hello"}],
"chat_template_kwargs": {"instruction": "Represent the query: "},
}
)

assert isinstance(request, EmbeddingChatInputRequest)
assert request.input == [{"role": "user", "content": "hello"}]
assert request.messages == [{"role": "user", "content": "hello"}]
assert request.chat_template_kwargs == {"instruction": "Represent the query: "}

def test_batched_input_messages_parses_as_batch_chat_input_request(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"input": [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
],
"chat_template_kwargs": {"instruction": "Represent the query: "},
}
)

assert isinstance(request, EmbeddingBatchChatInputRequest)
assert request.input == [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
]
assert request.messages == [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
]
assert request.chat_template_kwargs == {"instruction": "Represent the query: "}

def test_token_ids_still_parse_as_completion_request(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"input": [[1, 2, 3], [4, 5]],
}
)

assert isinstance(request, EmbeddingCompletionRequest)
assert request.input == [[1, 2, 3], [4, 5]]

def test_messages_still_parses_as_chat_request(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"messages": [{"role": "user", "content": "hello"}],
"chat_template_kwargs": {"instruction": "Represent the query: "},
}
)

assert isinstance(request, EmbeddingChatRequest)
assert request.messages == [{"role": "user", "content": "hello"}]
assert request.chat_template_kwargs == {"instruction": "Represent the query: "}

def test_batched_messages_parses_as_batch_chat_request(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"messages": [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
],
"chat_template_kwargs": {"instruction": "Represent the query: "},
}
)

assert isinstance(request, EmbeddingBatchChatRequest)
assert request.messages == [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
]
assert request.chat_template_kwargs == {"instruction": "Represent the query: "}


class TestResolveTruncation:
"""Unit tests for EmbedIOProcessor._resolve_cohere_truncation."""

Expand Down Expand Up @@ -324,3 +415,113 @@ def batch_render_chat(
},
)
]


class TestPreProcessOpenAIEmbeddingChatOnline:
"""Unit tests for OpenAI embedding chat preprocessing."""

class _FakeModelConfig:
max_model_len = 128
encoder_config: dict[str, object] = {}
pooler_config = None
multimodal_config = None
is_encoder_decoder = False

class _FakeRenderer:
tokenizer = object()

def __init__(self):
self.calls = []

def render_chat(
self,
all_messages,
chat_params,
tok_params,
prompt_extras=None,
):
self.calls.append(
{
"all_messages": all_messages,
"chat_params": chat_params,
"tok_params": tok_params,
"prompt_extras": prompt_extras,
}
)
return all_messages, [
{"prompt_token_ids": [index]} for index, _ in enumerate(all_messages)
]

@classmethod
def _make_handler(cls, renderer):
handler = object.__new__(EmbedIOProcessor)
handler.renderer = renderer
handler.model_config = cls._FakeModelConfig()
handler.chat_template = "template"
handler.chat_template_content_format = "auto"
handler.trust_request_chat_template = False
handler.enable_chunked_processing = False
return handler

@staticmethod
def _make_context(
request: (
EmbeddingChatRequest
| EmbeddingBatchChatRequest
| EmbeddingChatInputRequest
| EmbeddingBatchChatInputRequest
),
) -> PoolingServeContext[
EmbeddingChatRequest
| EmbeddingBatchChatRequest
| EmbeddingChatInputRequest
| EmbeddingBatchChatInputRequest
]:
return PoolingServeContext(
request=request,
pooling_params=PoolingParams(),
model_name="test",
request_id="embd-test",
)

def test_chat_template_kwargs_forwarded_for_batched_input_messages(self):
request = TypeAdapter(EmbeddingRequest).validate_python(
{
"model": "test",
"input": [
[{"role": "user", "content": "hello"}],
[{"role": "user", "content": "goodbye"}],
],
"add_generation_prompt": True,
"chat_template_kwargs": {"instruction": "Represent the query: "},
"mm_processor_kwargs": {"max_pixels": 1},
"cache_salt": "salt",
}
)
assert isinstance(request, EmbeddingBatchChatInputRequest)

renderer = self._FakeRenderer()
handler = self._make_handler(renderer)
ctx = self._make_context(request)

handler.pre_process_online(ctx)

assert ctx.engine_inputs == [
{"prompt_token_ids": [0]},
{"prompt_token_ids": [1]},
]
assert len(renderer.calls) == 1

call = renderer.calls[0]
assert call["all_messages"] == request.messages
assert call["prompt_extras"] == {
"mm_processor_kwargs": {"max_pixels": 1},
"cache_salt": "salt",
}

chat_template_kwargs = call["chat_params"].chat_template_kwargs
assert chat_template_kwargs["instruction"] == "Represent the query: "
assert chat_template_kwargs["add_generation_prompt"] is True
assert chat_template_kwargs["continue_final_message"] is False
assert "tools" not in chat_template_kwargs
assert chat_template_kwargs["tokenize"] is False
12 changes: 7 additions & 5 deletions vllm/entrypoints/pooling/base/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,11 +168,7 @@ class CompletionRequestMixin(OpenAIBaseModel):
# --8<-- [end:completion-extra-params]


class ChatRequestMixin(OpenAIBaseModel):
# --8<-- [start:chat-params]
messages: list[ChatCompletionMessageParam]
# --8<-- [end:chat-params]

class ChatRequestOptionsMixin(OpenAIBaseModel):
# --8<-- [start:chat-extra-params]
add_generation_prompt: bool = Field(
default=False,
Expand Down Expand Up @@ -256,6 +252,12 @@ def build_chat_params(
)


class ChatRequestMixin(ChatRequestOptionsMixin):
# --8<-- [start:chat-params]
messages: list[ChatCompletionMessageParam]
# --8<-- [end:chat-params]


class EncodingRequestMixin(OpenAIBaseModel):
# --8<-- [start:encoding-params]
encoding_format: EncodingFormat = "float"
Expand Down
77 changes: 77 additions & 0 deletions vllm/entrypoints/pooling/embed/io_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,9 @@
CohereEmbedContent,
CohereEmbedInput,
CohereEmbedRequest,
EmbeddingBatchChatInputRequest,
EmbeddingBatchChatRequest,
EmbeddingChatInputRequest,
EmbeddingChatRequest,
EmbeddingCompletionRequest,
)
Expand Down Expand Up @@ -66,6 +69,16 @@ def __init__(self, *args, **kwargs):
def pre_process_online(self, ctx: PoolingServeContext):
if isinstance(ctx.request, CohereEmbedRequest):
self._pre_process_cohere_online(ctx)
elif isinstance(
ctx.request,
(
EmbeddingChatRequest,
EmbeddingBatchChatRequest,
EmbeddingChatInputRequest,
EmbeddingBatchChatInputRequest,
),
):
self._pre_process_openai_chat_online(ctx)
else:
super().pre_process_online(ctx)

Expand Down Expand Up @@ -367,6 +380,70 @@ def create_pooling_params(self, request):
)
return super().create_pooling_params(request)

def _pre_process_openai_chat_online(
self,
ctx: PoolingServeContext[
EmbeddingChatRequest
| EmbeddingBatchChatRequest
| EmbeddingChatInputRequest
| EmbeddingBatchChatInputRequest
],
) -> None:
request = ctx.request
self._validate_chat_template(
request_chat_template=request.chat_template,
chat_template_kwargs=request.chat_template_kwargs,
trust_request_chat_template=self.trust_request_chat_template,
)

if isinstance(
request, (EmbeddingBatchChatRequest, EmbeddingBatchChatInputRequest)
):
all_messages = request.messages
else:
all_messages = [request.messages]
ctx.engine_inputs = self._batch_render_openai_chat(request, all_messages)

def _batch_render_openai_chat(
self,
request: (
EmbeddingChatRequest
| EmbeddingBatchChatRequest
| EmbeddingChatInputRequest
| EmbeddingBatchChatInputRequest
),
all_messages: Sequence[list[ChatCompletionMessageParam]],
) -> list[EngineInput]:
renderer = self.renderer
mm_config = self.model_config.multimodal_config

tok_params = request.build_tok_params(self.model_config)
chat_params = request.build_chat_params(
self.chat_template,
self.chat_template_content_format,
).with_defaults(
merge_kwargs(
None,
dict(
tools=None,
tokenize=is_mistral_tokenizer(renderer.tokenizer),
),
),
default_media_io_kwargs=(mm_config.media_io_kwargs if mm_config else None),
)

_, engine_inputs = renderer.render_chat(
all_messages,
chat_params,
tok_params,
prompt_extras={
k: v
for k in ("mm_processor_kwargs", "cache_salt")
if (v := getattr(request, k, None)) is not None
},
)
return engine_inputs

def _pre_process_cohere_online(self, ctx: PoolingServeContext) -> None:
"""Convert a ``CohereEmbedRequest`` into engine prompts.

Expand Down
Loading
Loading