57 lines
2.0 KiB
Python
57 lines
2.0 KiB
Python
#!/usr/bin/env python
|
|
from typing import Generator
|
|
import runpod
|
|
from utils import validate_and_convert_sampling_params, initialize_llm_engine, ServerlessConfig
|
|
from vllm.utils import random_uuid
|
|
|
|
serverless_config = ServerlessConfig()
|
|
llm, tokenizer = initialize_llm_engine()
|
|
|
|
def concurrency_modifier(current_concurrency) -> int:
|
|
return max(0, serverless_config.max_concurrency - current_concurrency)
|
|
|
|
async def handler(job: dict) -> Generator[dict, None, None]:
|
|
job_input = job["input"]
|
|
prompt = job_input.get("prompt")
|
|
apply_chat_template = job_input.get("apply_chat_template", False)
|
|
messages = job_input.get("messages")
|
|
|
|
if messages:
|
|
prompt = tokenizer.apply_chat_template(messages)
|
|
elif prompt and apply_chat_template:
|
|
prompt = tokenizer.apply_chat_template(prompt)
|
|
elif not prompt:
|
|
raise ValueError("Must specify prompt or messages")
|
|
|
|
streaming = job_input.get("streaming", False)
|
|
batch_size = job_input.get("batch_size", serverless_config.default_batch_size)
|
|
sampling_params = job_input.get("sampling_params", {})
|
|
|
|
validated_params = validate_and_convert_sampling_params(sampling_params)
|
|
request_id = random_uuid()
|
|
results_generator = llm.generate(prompt, validated_params, request_id)
|
|
|
|
batch, last_output_text = [], ""
|
|
async for request_output in results_generator:
|
|
for output in request_output.outputs:
|
|
usage = {"input": len(request_output.prompt_token_ids), "output": len(output.token_ids)}
|
|
|
|
if streaming:
|
|
batch.append({"text": output.text[len(last_output_text):], "usage": usage})
|
|
if len(batch) >= batch_size:
|
|
yield batch
|
|
batch = []
|
|
last_output_text = output.text
|
|
|
|
if not streaming:
|
|
yield [{"text": last_output_text, "usage": usage}]
|
|
|
|
if batch:
|
|
yield batch
|
|
|
|
runpod.serverless.start({
|
|
"handler": handler,
|
|
"concurrency_modifier": concurrency_modifier,
|
|
"return_aggregate_stream": True
|
|
})
|