Batched Tokens, cleanup
This commit is contained in:
+40
-20
@@ -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)),
|
||||
|
||||
+26
-4
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user