From bfa3703d548a0af1cd13b01438543739490132c1 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Wed, 13 Dec 2023 12:26:01 +0000 Subject: [PATCH] Concurrency modifier --- src/handler.py | 10 +++------- src/utils.py | 9 --------- 2 files changed, 3 insertions(+), 16 deletions(-) diff --git a/src/handler.py b/src/handler.py index b9f202d..6813a23 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,5 +1,4 @@ #!/usr/bin/env python - from typing import Generator import runpod from utils import validate_and_convert_sampling_params, initialize_llm_engine, JobManager, ServerlessConfig @@ -9,8 +8,8 @@ serverless_config = ServerlessConfig() job_manager = JobManager() llm = initialize_llm_engine() -def concurrency_modifier() -> int: - return max(0, serverless_config.max_concurrency - job_manager.total_running_jobs) +def concurrency_modifier(current_concurrency) -> int: + return max(0, serverless_config.max_concurrency - current_concurrency) async def handler(job: dict) -> Generator[dict, None, None]: job_input = job["input"] @@ -22,7 +21,6 @@ async def handler(job: dict) -> Generator[dict, None, None]: validated_params = validate_and_convert_sampling_params(sampling_params) request_id = random_uuid() results_generator = llm.generate(prompt, validated_params, request_id) - job_manager.increment_job_count() batch, last_output_text = [], "" async for request_output in results_generator: @@ -41,9 +39,7 @@ async def handler(job: dict) -> Generator[dict, None, None]: if batch: yield batch - - job_manager.decrement_job_count() - + runpod.serverless.start({ "handler": handler, "concurrency_modifier": concurrency_modifier, diff --git a/src/utils.py b/src/utils.py index 89974f5..6fa94ae 100644 --- a/src/utils.py +++ b/src/utils.py @@ -51,15 +51,6 @@ def initialize_llm_engine() -> AsyncLLMEngine: logging.error(f"Error initializing vLLM engine: {e}") raise -class JobManager: - def __init__(self): - self.total_running_jobs = 0 - - def increment_job_count(self): - self.total_running_jobs += 1 - - def decrement_job_count(self): - self.total_running_jobs -= 1 def validate_and_convert_sampling_params(params: Dict[str, Any]) -> SamplingParams: validated_params = {}