New Worker Stable

This commit is contained in:
alpayariyak
2023-12-29 08:21:09 +00:00
parent 2475cd7a66
commit a14c5cd388
4 changed files with 30 additions and 27 deletions
+25 -21
View File
@@ -5,12 +5,10 @@ from utils import validate_sampling_params, random_uuid
from engine import VLLMEngine
vllm_engine = VLLMEngine()
async def handler(job: dict) -> Generator[dict, None, None]:
job_input = job["input"]
llm_input, apply_chat_template = job_input.get(
"messages", job_input.get("prompt")
), job_input.get("apply_chat_template", False)
llm_input = job_input.get("messages", job_input.get("prompt"))
apply_chat_template = job_input.get("apply_chat_template", False)
if apply_chat_template or isinstance(llm_input, list):
llm_input = vllm_engine.tokenizer.apply_chat_template(llm_input)
@@ -25,34 +23,40 @@ async def handler(job: dict) -> Generator[dict, None, None]:
llm_input, validated_params, request_id
)
batch, last_output_text = [], ""
batch = {"tokens": []}
last_output_text = ""
n_input_tokens, is_first_output = 0, True
async for request_output in results_generator:
if is_first_output: # Count input tokens only once
n_input_tokens = len(request_output.prompt_token_ids)
is_first_output = False
for output in request_output.outputs:
usage = {
"input": len(request_output.prompt_token_ids),
"output": len(output.token_ids),
}
if stream:
batch.append(
{"text": output.text[len(last_output_text) :], "usage": usage}
batch["tokens"].append(
output.text[len(last_output_text):]
)
if len(batch) >= batch_size:
if len(batch["tokens"]) >= batch_size or request_output.finished:
batch["usage"] = {
"input": n_input_tokens,
"output": len(output.token_ids),
}
yield batch
batch = []
last_usage = batch["usage"]
batch = {"tokens": []}
last_output_text = output.text
if not stream:
yield [{"text": last_output_text, "usage": usage}]
if batch:
yield batch
yield {"tokens": [last_output_text],
"usage": last_usage}
runpod.serverless.start(
{
"handler": handler,
# "concurrency_modifier": vllm_engine.concurrency_modifier,
"concurrency_modifier": lambda x: vllm_engine.serverless_config.max_concurrency,
"return_aggregate_stream": True,
}