diff --git a/src/handler.py b/src/handler.py index fa8b6d3..db5fbb8 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,19 +1,16 @@ #!/usr/bin/env python -''' Contains the handler function that will be called by the serverless worker. ''' import os from typing import Generator - import runpod from utils import EngineConfig, validate_and_convert_sampling_params from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs, utils +# Default batch size, configurable via environment variable set in the Endpoint Template +DEFAULT_BATCH_SIZE = int(os.environ.get('DEFAULT_BATCH_SIZE', 10)) - -# Load the configuration +# Load the configuration for the vLLM engine config = EngineConfig() - -# Prepare the engine's arguments engine_args = AsyncEngineArgs( model=config.model_name, download_dir=config.model_base_path, @@ -22,35 +19,58 @@ engine_args = AsyncEngineArgs( dtype=config.dtype, disable_log_stats=config.disable_log_stats, quantization=config.quantization, - gpu_memory_utilization=0.97, + gpu_memory_utilization=0.98, ) -# Create the vLLM asynchronous engine +# Create the asynchronous vLLM engine llm = AsyncLLMEngine.from_engine_args(engine_args) - -# Handler function that will be called by the serverless worker async def handler(job: dict) -> Generator[str, None, None]: - job_input = job['input'] - prompt = job_input['prompt'] + """ + Asynchronous Generator Handler for the vLLM worker. + + Args: + job (dict): A dictionary containing job details, including the prompt and other parameters. + + Yields: + Generator[str, None, None]: A generator that yields generated text outputs. Format: List[str] + """ + # Extract the job inputs + job_input = job["input"] + prompt = job_input["prompt"] streaming = job_input.get("streaming", False) - sampling_params = job_input.get('sampling_params', {}) + batch_size = job_input.get("batch_size", DEFAULT_BATCH_SIZE) + sampling_params = job_input.get("sampling_params", {}) + + # Validate and convert sampling parameters validated_params = validate_and_convert_sampling_params(sampling_params) sampling_params_obj = SamplingParams(**validated_params) + # Generate a unique request ID request_id = utils.random_uuid() + + # Initialize the vLLM generator results_generator = llm.generate(prompt, sampling_params_obj, request_id) last_output_text = "" - + batch = [] + + # Process and yield the generated text async for request_output in results_generator: for output in request_output.outputs: - if output.text: - if streaming: - yield output.text[len(last_output_text):] - last_output_text = output.text - if not streaming: - yield last_output_text + if streaming: + batch.append(output.text[len(last_output_text):]) + if len(batch) >= batch_size: + yield batch + batch = [] + last_output_text = output.text + if not streaming: + yield [last_output_text] + + if batch and streaming: + yield batch + +# Start the serverless worker runpod.serverless.start({ "handler": handler, "concurrency_modifier": lambda _: int(os.environ.get('CONCURRENCY_MODIFIER', 100)), diff --git a/src/utils.py b/src/utils.py index 3373af2..fae52a3 100644 --- a/src/utils.py +++ b/src/utils.py @@ -1,15 +1,22 @@ import os class EngineConfig: + """ + Configuration for the vLLM engine. + """ def __init__(self, make_dirs=True): - self.model_name = os.getenv('MODEL_NAME', 'default_model') + self.model_name = os.getenv('MODEL_NAME') + if self.model_name is None: + raise ValueError("MODEL_NAME environment variable is not set") + self.tokenizer = os.getenv('TOKENIZER', self.model_name) self.model_base_path = os.getenv('MODEL_BASE_PATH', "/runpod-volume/") self.num_gpu_shard = int(os.getenv('NUM_GPU_SHARD', 1)) self.use_full_metrics = os.getenv('USE_FULL_METRICS', 'True') == 'True' - self.quantization = os.getenv('QUANTIZATION', None) - self.dtype = "auto" if str(self.quantization).lower() not in ['squeezellm', 'awq'] else "half" - self.disable_log_stats = os.getenv('DISABLE_LOG_STATS', 'False') == 'True' + self.quantization = str(os.getenv('QUANTIZATION', None)).lower() + self.quantization = self.quantization if self.quantization in ['squeezellm', 'awq'] else None + self.dtype = "auto" if self.quantization is None else "half" + self.disable_log_stats = os.getenv('DISABLE_LOG_STATS', 'True') == 'True' if make_dirs and not os.path.exists(self.model_base_path): os.makedirs(self.model_base_path) @@ -32,6 +39,14 @@ sampling_param_types = { # Function to convert sampling parameters to the right types def cast_sampling_param(value, target_type): + """ + Args: + value: The value to cast + target_type: The target type to cast to + + Returns: + The casted value if it can be casted, otherwise None + """ if value is None: return None try: @@ -42,6 +57,13 @@ def cast_sampling_param(value, target_type): # Function to validate and convert sampling parameters def validate_and_convert_sampling_params(sampling_params): + """ + Args: + sampling_params: The sampling parameters to validate and convert + + Returns: + The validated and converted sampling parameters + """ validated_params = {} for param_name, param_type in sampling_param_types.items(): param_value = sampling_params.get(param_name)