From a69ab1875dc0836f7a1326f50d91537652a7cf09 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Thu, 28 Dec 2023 23:06:49 +0000 Subject: [PATCH] vLLM job tracker, Refactor Concurrency Modifier, Serverless Config --- src/engine.py | 11 ++++++++++- src/handler.py | 12 +++--------- 2 files changed, 13 insertions(+), 10 deletions(-) diff --git a/src/engine.py b/src/engine.py index 5959e4f..3f81a2d 100644 --- a/src/engine.py +++ b/src/engine.py @@ -4,6 +4,7 @@ from typing import Union import torch from vllm import AsyncLLMEngine, AsyncEngineArgs from transformers import AutoTokenizer +from utils import ServerlessConfig class Tokenizer: @@ -28,6 +29,7 @@ class Tokenizer: class VLLMEngine: def __init__(self): self.config = self._initialize_config() + self.serverless_config = ServerlessConfig() self.tokenizer = Tokenizer(self.config["model"]) self.llm = self._initialize_llm() @@ -58,6 +60,13 @@ class VLLMEngine: logging.info("Using %s GPU shards", final_num_gpu_shard) return final_num_gpu_shard - def get_n_current_jobs(self): + def _get_n_current_jobs(self): total_sequences = len(self.llm.engine.scheduler.waiting) + len(self.llm.engine.scheduler.swapped) + len(self.llm.engine.scheduler.running) return total_sequences + + def concurrency_modifier(self, current_concurrency): + requested_concurrency = max(0, self.serverless_config.max_concurrency - self._get_n_current_jobs()) + if not self.config["disable_log_stats"]: + logging.info("Concurrency Modifier Requested Jobs: %s", requested_concurrency) + return requested_concurrency + diff --git a/src/handler.py b/src/handler.py index 9a067aa..0cb3881 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,17 +1,11 @@ #!/usr/bin/env python from typing import Generator import runpod -from utils import validate_sampling_params, ServerlessConfig, random_uuid +from utils import validate_sampling_params, random_uuid from engine import VLLMEngine - -serverless_config = ServerlessConfig() vllm_engine = VLLMEngine() - -def concurrency_modifier(current_concurrency) -> int: - return max(0, serverless_config.max_concurrency - vllm_engine.get_n_current_jobs) - async def handler(job: dict) -> Generator[dict, None, None]: job_input = job["input"] llm_input, apply_chat_template = job_input.get( @@ -22,7 +16,7 @@ async def handler(job: dict) -> Generator[dict, None, None]: llm_input = vllm_engine.tokenizer.apply_chat_template(llm_input) stream = job_input.get("stream", False) - batch_size = job_input.get("batch_size", serverless_config.default_batch_size) + batch_size = job_input.get("batch_size", vllm_engine.serverless_config.default_batch_size) sampling_params = job_input.get("sampling_params", {}) validated_params = validate_sampling_params(sampling_params) @@ -58,7 +52,7 @@ async def handler(job: dict) -> Generator[dict, None, None]: runpod.serverless.start( { "handler": handler, - "concurrency_modifier": concurrency_modifier, + "concurrency_modifier": vllm_engine.concurrency_modifier, "return_aggregate_stream": True, } )