diff --git a/.env b/.env deleted file mode 100644 index 505730d..0000000 --- a/.env +++ /dev/null @@ -1,4 +0,0 @@ -MODEL_NAME="mistralai/Mistral-7B-Instruct-v0.1" -MODEL_BASE_PATH="/devdisk/.cache/huggingface/hub" -DISABLE_LOG_STATS=1 -DISABLE_LOG_REQUESTS=1 \ No newline at end of file diff --git a/.gitignore b/.gitignore index 295faa0..b4e200e 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,5 @@ runpod.toml +*.pyc +.env +test/* \ No newline at end of file diff --git a/src/constants.py b/src/constants.py index e1866cc..6afc217 100644 --- a/src/constants.py +++ b/src/constants.py @@ -1,5 +1,5 @@ -DEFAULT_BATCH_SIZE = 10 -DEFAULT_MAX_CONCURRENCY = 100 +DEFAULT_BATCH_SIZE = 30 +DEFAULT_MAX_CONCURRENCY = 300 sampling_param_types = { "n": int, diff --git a/src/handler.py b/src/handler.py index 0119a72..f3d800b 100644 --- a/src/handler.py +++ b/src/handler.py @@ -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, }