update worker

This commit is contained in:
Jorg Doku
2023-07-14 11:17:25 -05:00
parent be65b7a99c
commit ab3f0e1c58
2 changed files with 26 additions and 6 deletions
+2 -2
View File
@@ -4,6 +4,6 @@
# Use the following additional index URL to get the latest PyTorch.
# --extra-index-url https://download.pytorch.org/whl/cu118
# vllm @ git+https://github.com/vllm-project/vllm.git@2b7d3aca2e1dd25fe26424f57c051af3b823cd71
# runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod
vllm==0.1.2
runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod
runpod @ git+https://github.com/runpod/runpod-python@multijob#egg=runpod
+24 -4
View File
@@ -1,14 +1,14 @@
#!/usr/bin/env python
''' Contains the handler function that will be called by the serverless. '''
from typing import Dict
import runpod
import runpod_vllm
# Start the VLLM serving layer on our RunPod worker.
from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs
from vllm.utils import random_uuid
# Prepare the model and tokenizer
MODEL = 'lmsys/vicuna-13b-v1.3'
MODEL = 'facebook/opt-125m'
TOKENIZER = 'hf-internal-testing/llama-tokenizer'
# Prepare the engine's arguments
@@ -23,6 +23,25 @@ engine_args = AsyncEngineArgs(
)
llm = AsyncLLMEngine.from_engine_args(engine_args)
def handler_fully_utilized() -> bool:
# Check VLLM metrics to see if we have reached maximum utilization. If we have, evaluate whether
# sleeping for X milliseconds will sustain the maximum utilization. If it does, sleep for
# X milliseconds and re-evaluate the check.
# A 7b model can process 5 iterations per second on A100. Assuming each iteration can handle
# up to 256 sequences, any sequences in waiting or swapped states will have to wait for at
# least one iteration before starting execution, which is around 1/5 a second.
#
# Sleeping for 200ms provides a sufficient delay for checking VLLM's queue state, even when
# using slower models such as 30B or higher. For models smaller than 7B, a smaller sleep
# delay of 20ms may be worth considering.
max_seq_per_iteration = 256
num_iters_threshold = 1
total_pending_sequences = len(llm.engine.scheduler.waiting) + len(llm.engine.scheduler.swapped)
# Check if we've surpassed the maximum number of sequences the vllm scheduler can handle per iteration.
return total_pending_sequences > max_seq_per_iteration * num_iters_threshold
# Validation
def validate_sampling_params(sampling_params):
def validate_int(value, default):
@@ -50,7 +69,7 @@ def validate_sampling_params(sampling_params):
top_p = validate_float(sampling_params.get('top_p'), 1.0)
top_k = validate_int(sampling_params.get('top_k'), -1)
use_beam_search = validate_bool(sampling_params.get('use_beam_search'), False)
stop = sampling_params.get('stop')
stop = sampling_params.get('stop', None)
ignore_eos = validate_bool(sampling_params.get('ignore_eos'), False)
max_tokens = validate_int(sampling_params.get('max_tokens'), 16)
logprobs = validate_float(sampling_params.get('logprobs'), None)
@@ -74,6 +93,7 @@ async def handler(job):
'''
This is the handler function that will be called by the serverless worker.
'''
print("Job received by handler: {}".format(job))
# Prompts
prompts = job['prompts']
@@ -99,4 +119,4 @@ async def handler(job):
ret = {"text": text_outputs}
return ret
runpod.serverless.start({"handler": handler})
runpod.serverless.start({"handler": handler, "handler_utilization": handler_fully_utilized})