Non-streaming OpenAI Chat Completions

This commit is contained in:
alpayariyak
2024-01-25 23:24:18 -05:00
parent 4cebe66b36
commit 9fc8e1e54c
3 changed files with 47 additions and 40 deletions
+1 -1
View File
@@ -151,7 +151,7 @@ You may either use a `prompt` or a list of `messages` as input. If you use `mess
|-----------------------|----------------------|--------------------|--------------------------------------------------------------------------------------------------------|
| `prompt` | str | | Prompt string to generate text based on. |
| `messages` | list[dict[str, str]] | | List of messages, which will automatically have the model's chat template applied. Overrides `prompt`. |
| `use_openai_format` | bool | False | Whether to return output in OpenAI format. `ALLOW_OPENAI_FORMAT` environment variable must be `1`, the input must be a `messages` list, and `stream` enabled. |
| `use_openai_format` | bool | False | Whether to return output in OpenAI format. `ALLOW_OPENAI_FORMAT` environment variable must be `1`, the input should preferably be a `messages` list, but `prompt` is accepted. |
| `apply_chat_template` | bool | False | Whether to apply the model's chat template to the `prompt`. |
| `sampling_params` | dict | {} | Sampling parameters to control the generation, like temperature, top_p, etc. |
| `stream` | bool | False | Whether to enable streaming of output. If True, responses are streamed as they are generated. |
+11 -7
View File
@@ -7,7 +7,7 @@ from vllm import AsyncLLMEngine, AsyncEngineArgs, SamplingParams
from vllm.entrypoints.openai.serving_chat import OpenAIServingChat
from vllm.entrypoints.openai.protocol import ChatCompletionRequest
from transformers import AutoTokenizer
from utils import count_physical_cores
from utils import count_physical_cores, DummyRequest
from constants import DEFAULT_MAX_CONCURRENCY
from dotenv import load_dotenv
@@ -106,20 +106,24 @@ class vLLMEngine:
async def generate_openai_chat(self, llm_input, validated_sampling_params, batch_size, stream, apply_chat_template, request_id: str) -> AsyncGenerator[dict, None]:
if not isinstance(llm_input, list):
raise ValueError("Input must be a list of messages")
if isinstance(llm_input, str):
llm_input = [{"role": "user", "content": llm_input}]
logging.warning("OpenAI Chat Completion format requires list input, converting to list and assigning 'user' role")
if not stream:
raise ValueError("OpenAI Chat Completion Format only supports streaming")
if not self.openai_engine:
raise ValueError("OpenAI Chat Completion format is disabled")
chat_completion_request = ChatCompletionRequest(
model=self.config["model"],
messages=llm_input,
stream=True,
stream=stream,
**validated_sampling_params,
)
response_generator = await self.openai_engine.create_chat_completion(chat_completion_request, None) # None for raw_request
response_generator = await self.openai_engine.create_chat_completion(chat_completion_request, DummyRequest())
if not stream:
yield json.loads(response_generator.model_dump_json())
else:
batch_contents = {}
batch_latest_choices = {}
batch_token_counter = 0
+3
View File
@@ -47,3 +47,6 @@ class JobInput:
self.validated_sampling_params = validate_sampling_params(job.get("sampling_params", {}))
self.request_id = random_uuid()
class DummyRequest:
async def is_disconnected(self):
return False