Concurrency and Prompt Template fix

This commit is contained in:
alpayariyak
2023-12-29 00:42:40 +00:00
parent a69ab1875d
commit ca8b02e392
2 changed files with 12 additions and 7 deletions
+10 -6
View File
@@ -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
+2 -1
View File
@@ -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,
}
)