This commit is contained in:
Justin Merrell
2023-11-29 21:32:40 -05:00
parent 820a21f32c
commit 9b9c1cf503
7 changed files with 75 additions and 426 deletions
+64 -242
View File
@@ -1,50 +1,31 @@
#!/usr/bin/env python
''' Contains the handler function that will be called by the serverless worker. '''
# Start the vLLM serving layer on our RunPod worker.
from typing import Generator
from metrics import vllm_log_system_stats
from templates import DEFAULT_TEMPLATE, LLAMA2_TEMPLATE
from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs
from vllm.utils import random_uuid
import runpod
import os
from typing import Generator
import runpod
from metrics import vllm_log_system_stats
from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs, utils
NUM_GPU_SHARD = int(os.environ.get('NUM_GPU_SHARD', 1)) # Number of GPUs to shard the model across
# Prepare the model and tokenizer
MODEL_NAME = os.environ.get('MODEL_NAME')
STREAMING = os.environ.get('STREAMING', False) == 'True'
TOKENIZER = os.environ.get('TOKENIZER', None)
DTYPE = "auto"
MODEL_NAME = os.environ["MODEL_NAME"]
TOKENIZER = os.environ.get('TOKENIZER', MODEL_NAME)
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH')
if not os.path.exists(BASE_VOLUME):
os.makedirs(BASE_VOLUME)
USE_FULL_METRICS = os.environ.get('USE_FULL_METRICS', True)
MAX_CONCURRENCY = os.environ.get('MAX_CONCURRENCY', 200)
TOTAL_RUNNING_JOBS = 0
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH', "/runpod-volume/")
if not os.path.exists(MODEL_BASE_PATH):
os.makedirs(MODEL_BASE_PATH)
USE_FULL_METRICS = os.environ.get('USE_FULL_METRICS', True) # From the SDK, need to review later.
# Set up quantization-related parameters
QUANTIZATION = os.environ.get('QUANTIZATION', None)
if type(QUANTIZATION) is str and QUANTIZATION.lower() in ["awq", "squeezellm"]:
QUANTIZATION = None
print("Invalid quantization parameter. Using default value of None.")
else:
DTYPE = "half"
if not MODEL_NAME:
print("Error: The model has not been provided.")
if not TOKENIZER or len(TOKENIZER) == 0:
print("Error: The tokenizer has not been provided. Defaulting to MODEL_NAME.")
# Tensor parallelism
try:
NUM_GPU_SHARD = int(os.environ.get('NUM_GPU_SHARD', 1))
except ValueError:
print("Error: NUM_GPU_SHARD should be an integer. Using default value of 1.")
NUM_GPU_SHARD = 1
DTYPE = "auto" if str(QUANTIZATION).lower() not in ['squeezellm', 'awq'] else "half"
# Prepare the engine's arguments
@@ -67,11 +48,6 @@ llm.engine._log_system_stats = lambda x, y: vllm_log_system_stats(
llm.engine, x, y)
def concurrency_controller() -> int:
global TOTAL_RUNNING_JOBS
return max(0, MAX_CONCURRENCY - TOTAL_RUNNING_JOBS)
def prepare_metrics() -> dict:
# The vLLM metrics are updated every 5 seconds, see metrics.py for the _LOGGING_INTERVAL_SEC field.
if hasattr(llm.engine, 'metrics'):
@@ -147,232 +123,78 @@ async def handler_streaming(job: dict) -> Generator[dict[str, list], None, None]
'''
This is the handler function that will be called by the serverless worker.
'''
print("Job received by handler: {}".format(job))
global TOTAL_RUNNING_JOBS
# Retrieve the job input.
print(f"Job received by handler: {job}")
job_input = job['input']
# Utilize the built-in llama2 template if a llama2 base model is being employed.
llama_models = ["llama-2-7b-chat-hf", "llama-2-13b-chat-hf", "llama-2-70b-chat-hf", "elinas/chronos-13b-v2"]
# Create the prompt using the template.
prompt = job_input['prompt']
sampling_params = job_input.get('sampling_params', None)
# Validate and set sampling parameters
sampling_params = validate_and_set_sampling_params(job_input.get('sampling_params', None))
# Print job input and sampling parameters
print("Job Input:", job_input)
print("Sampling Parameters:", sampling_params)
# Might be able to remove this later
sampling_params = validate_and_set_sampling_params(sampling_params)
# Send request to VLLM
request_id = random_uuid()
TOTAL_RUNNING_JOBS += 1
request_id = utils.random_uuid()
results_generator = llm.generate(prompt, sampling_params, request_id)
# Keep track of the stream's information to perform the appropriate chunking.
class Tracker():
def __init__(self):
self.positions = None
self.stream_index = 0
def inc_stream_idx(self):
self.stream_index += 1
tracker = Tracker()
def extract_next_chunk(request_output):
"""
Extracts and processes generated chunks and token counts from the request output.
Args:
request_output (CompletionOutput): The output of a language model request.
Returns:
tuple: A tuple containing two lists - chunk_outputs (extracted chunks) and num_output_tokens (generated token counts).
"""
chunk_outputs = [] # List to store extracted chunks
num_output_tokens = [] # List to store generated token counts
# Iterate over each completion in the request output
for idx, completion in enumerate(request_output.outputs):
# Extract the current chunk position from the tracker
chunk_pos = tracker.positions[idx]['chunk_pos']
# Append the chunk to the output
chunk_outputs.append(completion.text[chunk_pos:])
# Update the chunk position in the tracker
tracker.positions[idx]['chunk_pos'] = len(completion.text)
# Calculate the number of generated tokens in the current completion
num_generated_tokens = len(completion.token_ids) - tracker.positions[idx]['token_pos']
# Append the token count to the output
num_output_tokens.append(num_generated_tokens)
# Update the token position in the tracker
tracker.positions[idx]['token_pos'] = len(completion.token_ids)
return chunk_outputs, num_output_tokens
stream_index = 0
chunk_positions = None
aggregate_text = []
aggregate_metrics = {'input_tokens': 0, 'output_tokens': 0}
async for request_output in results_generator:
# Initialize chunk positions if not already done
if tracker.positions is None:
tracker.positions = [{'chunk_pos': 0, 'token_pos': 0}] * len(request_output.outputs)
if chunk_positions is None:
chunk_positions = [0] * len(request_output.outputs)
text_outputs, output_tokens = [], []
for idx, output in enumerate(request_output.outputs):
chunk_pos = chunk_positions[idx]
text_outputs.append(output.text[chunk_pos:])
chunk_positions[idx] = len(output.text)
output_tokens.append(len(output.token_ids) - chunk_pos)
# Metrics for the vLLM serverless worker
runpod_metrics = prepare_metrics() if USE_FULL_METRICS else {}
# Number of generated sequences
num_seqs = sampling_params.n
# Extract the next chunk from the output
text_outputs, output_tokens = extract_next_chunk(request_output)
# Record job input and token counts
# runpod_metrics['job_input'] = job_input
# Only include the input_tokens count for the very first stream response. This is to avoid duplicate counting.
if tracker.stream_index == 0:
if stream_index == 0:
input_tokens_count = len(request_output.prompt_token_ids)
runpod_metrics['input_tokens'] = sum([input_tokens_count] * num_seqs)
else:
runpod_metrics['input_tokens'] = sum([0] * num_seqs)
input_tokens_count = 0
# Include the output tokens count [#, #, #, ...]
runpod_metrics['output_tokens'] = sum(output_tokens)
runpod_metrics.update({
"input_tokens": input_tokens_count,
"output_tokens": sum(output_tokens),
"scenario": 'stream',
"stream_index": stream_index
})
# Store the scenario type and stream index
runpod_metrics['scenario'] = 'stream'
runpod_metrics['stream_index'] = tracker.stream_index
stream_index += 1
# Increment the index within the stream
tracker.inc_stream_idx()
ret = {
yield {
"text": text_outputs,
"input_tokens": runpod_metrics['input_tokens'],
"output_tokens": runpod_metrics['output_tokens']
}
# Include metrics for the job.
runpod.serverless.modules.rp_metrics.metrics_collector.push_metrics_internal(
job_id=job['id'],
metrics=runpod_metrics
)
# Aggregate text and metrics
if not aggregate_text:
aggregate_text = [""] * len(text_outputs)
for idx, text in enumerate(text_outputs):
aggregate_text[idx] += text
aggregate_metrics['input_tokens'] += runpod_metrics['input_tokens']
aggregate_metrics['output_tokens'] += runpod_metrics['output_tokens']
# Keep track of the final output
final_output = request_output
# Include metrics in the highest level for the job output for aggregrate.
def aggregate_function(streamed_outputs):
aggregate_output = [""] * len(streamed_outputs[0]['text'])
for stream in streamed_outputs:
for id, seq in enumerate(stream['text']):
aggregate_output[id] += seq
# Number of generated sequences
num_seqs = sampling_params.n
# Aggregate metrics to expose to the user
input_tokens = len(final_output.prompt_token_ids) * num_seqs
output_tokens = sum([len(output.token_ids) for output in final_output.outputs])
return {
"text": aggregate_output,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
}
# Update the aggregate transformation function
runpod.serverless.modules.rp_metrics.metrics_collector.update_stream_aggregate(
job_id=job['id'],
aggregate_function=aggregate_function
)
# Yield the output
yield ret
TOTAL_RUNNING_JOBS -= 1
async def handler(job: dict) -> dict[str, list]:
'''
This is the handler function that will be called by the serverless worker.
'''
print("Job received by handler: {}".format(job))
global TOTAL_RUNNING_JOBS
# Retrieve the job input.
job_input = job['input']
# Create the prompt using the template.
prompt = job_input['prompt']
# Validate and set sampling parameters
sampling_params = validate_and_set_sampling_params(job_input.get('sampling_params', None))
# Print job input and sampling parameters
print("Job Input:", job_input)
print("Sampling Parameters:", sampling_params)
# Send request to VLLM
request_id = random_uuid()
TOTAL_RUNNING_JOBS += 1
results_generator = llm.generate(prompt, sampling_params, request_id)
# Get the final generated output
final_output = None
async for request_output in results_generator:
final_output = request_output
# Extract prompt and text outputs
prompt = final_output.prompt
text_outputs = [output.text for output in final_output.outputs]
# Number of generated sequences
num_seqs = sampling_params.n
# Prepare metrics if full metrics are enabled
runpod_metrics = prepare_metrics() if USE_FULL_METRICS else {}
# Record job input and token counts
# runpod_metrics['job_input'] = job_input
runpod_metrics['input_tokens'] = len(final_output.prompt_token_ids) * num_seqs
runpod_metrics['output_tokens'] = sum([len(output.token_ids) for output in final_output.outputs])
# Store the scenario type
runpod_metrics['scenario'] = 'batch'
# Include metrics for the job.
runpod.serverless.modules.rp_metrics.metrics_collector.push_metrics_internal(
job_id=job['id'],
metrics=runpod_metrics
)
ret = {
"text": text_outputs,
"input_tokens": runpod_metrics['input_tokens'],
"output_tokens": runpod_metrics['output_tokens']
yield {
"text": aggregate_text,
"input_tokens": aggregate_metrics['input_tokens'],
"output_tokens": aggregate_metrics['output_tokens']
}
TOTAL_RUNNING_JOBS -= 1
return ret
def concurrency_modifier() -> int:
return 100
# Start the serverless worker with appropriate settings
if STREAMING:
print("Starting the vLLM serverless worker with streaming enabled.")
runpod.serverless.start({
"handler": handler_streaming,
"concurrency_controller": concurrency_controller,
"return_aggregate_stream": True
})
else:
print("Starting the vLLM serverless worker with streaming disabled.")
runpod.serverless.start({
"handler": handler,
"concurrency_controller":
concurrency_controller
})
runpod.serverless.start({
"handler": handler_streaming,
"concurrency_modifier": concurrency_modifier,
"return_aggregate_stream": True
})