From ef3c3037437895056e0c395196908f7362bdfdb3 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Fri, 19 Jan 2024 15:34:25 +0000 Subject: [PATCH] Added support for `n` parameter --- src/handler.py | 44 ++++++++++++++++++++++++++++---------------- 1 file changed, 28 insertions(+), 16 deletions(-) diff --git a/src/handler.py b/src/handler.py index fafe913..7ce48d8 100644 --- a/src/handler.py +++ b/src/handler.py @@ -11,19 +11,22 @@ async def handler(job: dict) -> Generator[dict, None, None]: job_input = job["input"] llm_input = job_input.get("messages", job_input.get("prompt")) if job_input.get("apply_chat_template", False) or isinstance(llm_input, list): - llm_input = vllm_engine.tokenizer.apply_chat_template(llm_input) - + llm_input = vllm_engine.tokenizer.apply_chat_template(llm_input) + stream = job_input.get("stream", False) batch_size = job_input.get("batch_size", vllm_engine.serverless_config.batch_size) + validated_params = validate_sampling_params(job_input.get("sampling_params", {})) - request_id = random_uuid() - results_generator = vllm_engine.llm.generate( - llm_input, validated_params, request_id + llm_input, validated_params, random_uuid() ) - - batch, last_output_text, n_input_tokens, is_first_output = {"tokens": []}, "", 0, True - + + n_responses, n_input_tokens, is_first_output = validated_params.n, 0, True + last_output_texts, token_counters= ["" for _ in range(n_responses)], {"batch": 0, "total": 0} + + batch = { + "choices": [{"tokens": []} for _ in range(n_responses)], + } async for request_output in results_generator: if is_first_output: # Count input tokens only once @@ -31,24 +34,33 @@ async def handler(job: dict) -> Generator[dict, None, None]: is_first_output = False for output in request_output.outputs: + output_index = output.index + token_counters["total"] += 1 if stream: - batch["tokens"].append(output.text[len(last_output_text):]) + new_output = output.text[len(last_output_texts[output_index]):] + batch["choices"][output_index]["tokens"].append(new_output) + token_counters["batch"] += 1 - if len(batch["tokens"]) >= batch_size: + if token_counters["batch"] >= batch_size: batch["usage"] = { "input": n_input_tokens, - "output": len(output.token_ids), + "output": token_counters["total"], } yield batch - batch = {"tokens": []} + batch = { + "choices": [{"tokens": []} for _ in range(n_responses)], + } + token_counters["batch"] = 0 - last_output_text = output.text + last_output_texts[output_index] = output.text if not stream: - batch["tokens"].append(last_output_text) + for output_index, output in enumerate(last_output_texts): + batch["choices"][output_index]["tokens"] = [output] + token_counters["batch"] += 1 - if len(batch["tokens"]) > 0: - batch["usage"] = {"input": n_input_tokens, "output": len(output.token_ids)} + if token_counters["batch"] > 0: + batch["usage"] = {"input": n_input_tokens, "output": token_counters["total"]} yield batch runpod.serverless.start(