From ca8b02e3926c8d587aa6ccd11d864b9d10af62dd Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Fri, 29 Dec 2023 00:42:40 +0000 Subject: [PATCH] Concurrency and Prompt Template fix --- src/engine.py | 16 ++++++++++------ src/handler.py | 3 ++- 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/src/engine.py b/src/engine.py index 3f81a2d..c270ac5 100644 --- a/src/engine.py +++ b/src/engine.py @@ -13,14 +13,16 @@ class Tokenizer: self.has_chat_template = bool(self.tokenizer.chat_template) def apply_chat_template(self, input: Union[str, list[dict[str, str]]]) -> str: - if isinstance(input, list) and not self.has_chat_template: - raise ValueError( - "Chat template does not exist for this model, you must provide a single string input instead of a list of messages" - ) + if isinstance(input, list): + if not self.has_chat_template: + raise ValueError( + "Chat template does not exist for this model, you must provide a single string input instead of a list of messages" + ) elif isinstance(input, str): input = [{"role": "user", "content": input}] else: raise ValueError("Input must be a string or a list of messages") + return self.tokenizer.apply_chat_template( input, tokenize=False, add_generation_prompt=True ) @@ -39,7 +41,7 @@ class VLLMEngine: "download_dir": os.getenv("MODEL_BASE_PATH", "/runpod-volume/"), "quantization": os.getenv("QUANTIZATION"), "dtype": "auto" if os.getenv("QUANTIZATION") is None else "half", - "disable_log_stats": bool(int(os.getenv("DISABLE_LOG_STATS", 1))), + "disable_log_stats": bool(int(os.getenv("DISABLE_LOG_STATS", 0))), "gpu_memory_utilization": float(os.getenv("GPU_MEMORY_UTILIZATION", 0.98)), "tensor_parallel_size": self._get_num_gpu_shard(), } @@ -65,8 +67,10 @@ class VLLMEngine: return total_sequences def concurrency_modifier(self, current_concurrency): - requested_concurrency = max(0, self.serverless_config.max_concurrency - self._get_n_current_jobs()) + n_current_jobs = self._get_n_current_jobs() + requested_concurrency = max(0, self.serverless_config.max_concurrency - n_current_jobs) if not self.config["disable_log_stats"]: + logging.info("Current Jobs: %s", n_current_jobs) logging.info("Concurrency Modifier Requested Jobs: %s", requested_concurrency) return requested_concurrency diff --git a/src/handler.py b/src/handler.py index 0cb3881..235b3cf 100644 --- a/src/handler.py +++ b/src/handler.py @@ -52,7 +52,8 @@ async def handler(job: dict) -> Generator[dict, None, None]: runpod.serverless.start( { "handler": handler, - "concurrency_modifier": vllm_engine.concurrency_modifier, + # "concurrency_modifier": vllm_engine.concurrency_modifier, + "concurrency_modifier": lambda x: vllm_engine.serverless_config.max_concurrency, "return_aggregate_stream": True, } )