diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index e4c3040c9..75208874f 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -194,7 +194,8 @@ class EmbeddingReqInput: if is_single: if self.rid is None: self.rid = uuid.uuid4().hex - self.sampling_params = {"max_new_tokens": 0} + if self.sampling_params is None: + self.sampling_params = {"max_new_tokens": 1} else: # support select operation self.batch_size = ( @@ -205,9 +206,10 @@ class EmbeddingReqInput: else: if not isinstance(self.rid, list): raise ValueError("The rid should be a list.") - self.sampling_params = [ - {"max_new_tokens": 0} for _ in range(self.batch_size) - ] + if self.sampling_params is None: + self.sampling_params = [ + {"max_new_tokens": 1} for _ in range(self.batch_size) + ] @dataclass diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 8711c127d..43c70ac7c 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -262,6 +262,7 @@ class TokenizerManager: ): yield response else: + assert self.is_generation await self._wait_for_cache_prefill_response(event, state, obj, rid, request) yield input_ids diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 77941c8af..a7a78bde3 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -499,6 +499,8 @@ class ModelTpServer: req.embedding = embeddings[i] if req is not self.current_inflight_req: # Inflight reqs' prefill is not finished + # dummy output token for embedding models + req.output_ids.append(0) req.check_finished() if req.finished(): diff --git a/python/sglang/srt/openai_api/adapter.py b/python/sglang/srt/openai_api/adapter.py index c12138391..5c4b9719d 100644 --- a/python/sglang/srt/openai_api/adapter.py +++ b/python/sglang/srt/openai_api/adapter.py @@ -34,7 +34,7 @@ from sglang.srt.conversation import ( generate_chat_conv, register_conv_template, ) -from sglang.srt.managers.io_struct import GenerateReqInput +from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.openai_api.protocol import ( BatchRequest, BatchResponse, @@ -52,6 +52,7 @@ from sglang.srt.openai_api.protocol import ( CompletionResponseStreamChoice, CompletionStreamResponse, DeltaMessage, + EmbeddingObject, EmbeddingRequest, EmbeddingResponse, ErrorResponse, @@ -1016,10 +1017,10 @@ async def v1_chat_completions(tokenizer_manager, raw_request: Request): def v1_embedding_request(all_requests, tokenizer_manager): prompts = [] sampling_params_list = [] - first_prompt_type = type(all_requests[0].prompt) + first_prompt_type = type(all_requests[0].input) for request in all_requests: - prompt = request.prompt + prompt = request.input assert ( type(prompt) == first_prompt_type ), "All prompts must be of the same type in file input settings" @@ -1046,17 +1047,26 @@ def v1_embedding_request(all_requests, tokenizer_manager): return adapted_request, all_requests -def v1_embedding_response(request, ret, to_file=False): - response = [] +def v1_embedding_response(ret, model_path, to_file=False): + embedding_objects = [] + prompt_tokens = 0 for idx, ret_item in enumerate(ret): - response.append( - EmbeddingResponse( + embedding_objects.append( + EmbeddingObject( + embedding=ret[idx]["embedding"], index=idx, - embedding=ret[idx], - object="embedding", ) ) - return response + prompt_tokens += ret[idx]["meta_info"]["prompt_tokens"] + + return EmbeddingResponse( + data=embedding_objects, + model=model_path, + usage=UsageInfo( + prompt_tokens=prompt_tokens, + total_tokens=prompt_tokens, + ), + ) async def v1_embeddings(tokenizer_manager, raw_request: Request): @@ -1074,7 +1084,7 @@ async def v1_embeddings(tokenizer_manager, raw_request: Request): if not isinstance(ret, list): ret = [ret] - response = v1_embedding_response(request, ret) + response = v1_embedding_response(ret, tokenizer_manager.model_path) return response diff --git a/python/sglang/srt/openai_api/protocol.py b/python/sglang/srt/openai_api/protocol.py index 75f0a1aab..758e48ede 100644 --- a/python/sglang/srt/openai_api/protocol.py +++ b/python/sglang/srt/openai_api/protocol.py @@ -319,8 +319,14 @@ class EmbeddingRequest(BaseModel): user: Optional[str] = None -class EmbeddingResponse(BaseModel): - index: str - embedding: List[float] = None +class EmbeddingObject(BaseModel): + embedding: List[float] + index: int object: str = "embedding" + + +class EmbeddingResponse(BaseModel): + data: List[EmbeddingObject] + model: str + object: str = "list" usage: Optional[UsageInfo] = None diff --git a/python/sglang/srt/server.py b/python/sglang/srt/server.py index d6e3f31ec..ed611242f 100644 --- a/python/sglang/srt/server.py +++ b/python/sglang/srt/server.py @@ -60,6 +60,7 @@ from sglang.srt.openai_api.adapter import ( v1_chat_completions, v1_completions, v1_delete_file, + v1_embeddings, v1_files_create, v1_retrieve_batch, v1_retrieve_file, @@ -176,6 +177,12 @@ async def openai_v1_chat_completions(raw_request: Request): return await v1_chat_completions(tokenizer_manager, raw_request) +@app.post("/v1/embeddings") +async def openai_v1_embeddings(raw_request: Request): + response = await v1_embeddings(tokenizer_manager, raw_request) + return response + + @app.get("/v1/models") def available_models(): """Show available models.""" @@ -412,7 +419,7 @@ def _wait_and_warmup(server_args, pipe_finish_writer): # Send a warmup request request_name = "/generate" if model_info["is_generation"] else "/encode" - max_new_tokens = 8 if model_info["is_generation"] else 0 + max_new_tokens = 8 if model_info["is_generation"] else 1 try: for _ in range(server_args.dp_size): res = requests.post( diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index d5051ffc1..2bc37b682 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -6,6 +6,7 @@ from sglang.test.test_utils import run_unittest_files suites = { "minimal": [ "test_eval_accuracy.py", + "test_embedding_openai_server.py", "test_openai_server.py", "test_vision_openai_server.py", "test_chunked_prefill.py", diff --git a/test/srt/test_embedding_openai_server.py b/test/srt/test_embedding_openai_server.py new file mode 100644 index 000000000..d60ae5068 --- /dev/null +++ b/test/srt/test_embedding_openai_server.py @@ -0,0 +1,87 @@ +import json +import time +import unittest + +import openai + +from sglang.srt.hf_transformers_utils import get_tokenizer +from sglang.srt.openai_api.protocol import EmbeddingObject +from sglang.srt.utils import kill_child_process +from sglang.test.test_utils import popen_launch_server + + +class TestOpenAIServer(unittest.TestCase): + + @classmethod + def setUpClass(cls): + cls.model = "intfloat/e5-mistral-7b-instruct" + cls.base_url = "http://127.0.0.1:8157" + cls.api_key = "sk-123456" + cls.process = popen_launch_server( + cls.model, cls.base_url, timeout=300, api_key=cls.api_key + ) + cls.base_url += "/v1" + cls.tokenizer = get_tokenizer(cls.model) + + @classmethod + def tearDownClass(cls): + kill_child_process(cls.process.pid) + + def run_embedding(self, use_list_input, token_input): + client = openai.Client(api_key=self.api_key, base_url=self.base_url) + prompt = "The capital of France is" + if token_input: + prompt_input = self.tokenizer.encode(prompt) + num_prompt_tokens = len(prompt_input) + else: + prompt_input = prompt + num_prompt_tokens = len(self.tokenizer.encode(prompt)) + + if use_list_input: + prompt_arg = [prompt_input, prompt_input] + num_prompts = len(prompt_arg) + else: + prompt_arg = prompt_input + num_prompts = 1 + + response = client.embeddings.create( + input=prompt_arg, + model=self.model, + ) + + assert len(response.data) == num_prompts + assert isinstance(response.data, list) + assert response.data[0].embedding + assert response.data[0].index is not None + assert response.data[0].object == "embedding" + assert response.model == self.model + assert response.object == "list" + assert ( + response.usage.prompt_tokens == num_prompt_tokens + ), f"{response.usage.prompt_tokens} vs {num_prompt_tokens}" + assert ( + response.usage.total_tokens == num_prompt_tokens + ), f"{response.usage.total_tokens} vs {num_prompt_tokens}" + + def run_batch(self): + # FIXME not implemented + pass + + def test_embedding(self): + # TODO the fields of encoding_format, dimensions, user are skipped + # TODO support use_list_input + for use_list_input in [False]: + for token_input in [False, True]: + self.run_embedding(use_list_input, token_input) + + def test_batch(self): + self.run_batch() + + +if __name__ == "__main__": + unittest.main(warnings="ignore") + + # t = TestOpenAIServer() + # t.setUpClass() + # t.test_embedding() + # t.tearDownClass()