latest version

This commit is contained in:
Jorg Doku
2023-08-03 10:05:51 -05:00
parent 728e4f4f2e
commit 401a4637fa
4 changed files with 106 additions and 55 deletions
+16 -7
View File
@@ -31,22 +31,31 @@ RUN --mount=type=cache,target=/root/.cache/pip \
ADD src .
# Quick temporary updates
RUN pip install git+https://github.com/runpod/runpod-python@main#egg=runpod --compile
RUN pip install git+https://github.com/runpod/runpod-python@async_gen_test#egg=runpod --compile
# Prepare the models inside the docker image
ARG HUGGING_FACE_HUB_TOKEN=NONE
ARG HUGGING_FACE_HUB_TOKEN=
ENV HUGGING_FACE_HUB_TOKEN=$HUGGING_FACE_HUB_TOKEN
# Prepare argument for the model and tokenizer
ARG MODEL=
ARG MODEL_NAME=""
ENV MODEL_NAME=$MODEL_NAME
ARG MODEL_REVISION="main"
ENV MODEL_REVISION=$MODEL_REVISION
ARG MODEL_BASE_PATH="/runpod-volume/"
ENV MODEL_BASE_PATH=$MODEL_BASE_PATH
ARG TOKENIZER=
ENV MODEL=$MODEL
ENV TOKENIZER=$TOKENIZER
ARG STREAMING=
ENV STREAMING=$STREAMING
ENV HF_DATASETS_CACHE="/runpod-volume/huggingface-cache/datasets"
ENV HUGGINGFACE_HUB_CACHE="/runpod-volume/huggingface-cache/hub"
ENV TRANSFORMERS_CACHE="/runpod-volume/huggingface-cache/hub"
# Download the models
RUN mkdir -p /model
RUN MODEL=$MODEL HUGGING_FACE_HUB_TOKEN=$HUGGING_FACE_HUB_TOKEN python -u /download_model.py
RUN MODEL_NAME=$MODEL_NAME MODEL_REVISION=$MODEL_REVISION MODEL_BASE_PATH=$MODEL_BASE_PATH HUGGING_FACE_HUB_TOKEN=$HUGGING_FACE_HUB_TOKEN python -u /download_model.py
# Start the handler
CMD MODEL=$MODEL TOKENIZER=$TOKENIZER python -u /handler.py
CMD STREAMING=$STREAMING MODEL_NAME=$MODEL_NAME MODEL_BASE_PATH=$MODEL_BASE_PATH TOKENIZER=$TOKENIZER python -u /handler.py
+2 -1
View File
@@ -1,5 +1,6 @@
# Required Python packages get listed here, one per line.
# Recomended to lock the version number to avoid unexpected changes.
fastapi==0.99.1
vllm==0.1.2
huggingface-hub==0.16.4
runpod @ git+https://github.com/runpod/runpod-python@main#egg=runpod
runpod @ git+https://github.com/runpod/runpod-python@async_gen_test#egg=runpod
+14 -7
View File
@@ -3,12 +3,19 @@ from huggingface_hub import snapshot_download
# Get the hugging face token
HUGGING_FACE_HUB_TOKEN = os.environ.get('HUGGING_FACE_HUB_TOKEN', None)
MODEL = os.environ.get('MODEL', None)
MODEL_NAME = os.environ.get('MODEL_NAME')
MODEL_REVISION = os.environ.get('MODEL_REVISION', "main")
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH', '/runpod-volume/')
# Download the model from hugging face
if HUGGING_FACE_HUB_TOKEN and MODEL:
snapshot_download(
MODEL,
local_dir="/model/{}".format(MODEL),
token=HUGGING_FACE_HUB_TOKEN
)
download_kwargs = {}
if HUGGING_FACE_HUB_TOKEN:
download_kwargs["token"] = HUGGING_FACE_HUB_TOKEN
snapshot_download(
MODEL_NAME,
revision=MODEL_REVISION,
local_dir=f"{MODEL_BASE_PATH}{MODEL_NAME.split('/')[1]}",
**download_kwargs
)
+74 -40
View File
@@ -9,18 +9,17 @@ import runpod
import os
# Prepare the model and tokenizer
MODEL = os.environ.get('MODEL', None)
MODEL_NAME = os.environ.get('MODEL_NAME')
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH', '/runpod-volume/')
STREAMING = os.environ.get('STREAMING', False)
TOKENIZER = os.environ.get('TOKENIZER', None)
if not MODEL:
if not MODEL_NAME:
print("Error: The model has not been provided.")
if not TOKENIZER:
print("Error: The tokenizer has not been provided.")
# Prepare the engine's arguments
engine_args = AsyncEngineArgs(
model="/model/{}".format(MODEL),
model=f"{MODEL_BASE_PATH}{MODEL_NAME.split('/')[1]}",
tokenizer=TOKENIZER,
tokenizer_mode="auto",
tensor_parallel_size=1,
@@ -72,7 +71,7 @@ def validate_sampling_params(sampling_params):
sampling_params.get('use_beam_search'), False)
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)
max_tokens = validate_int(sampling_params.get('max_tokens'), 256)
logprobs = validate_float(sampling_params.get('logprobs'), None)
return {
@@ -91,7 +90,7 @@ def validate_sampling_params(sampling_params):
}
async def handler(job):
async def handler_streaming(job):
'''
This is the handler function that will be called by the serverless worker.
'''
@@ -101,16 +100,13 @@ async def handler(job):
job_input = job['input']
# Prompts
if MODEL == "Llama-2-7b-chat-hf" or MODEL == "Llama-2-13b-chat-hf":
if MODEL_NAME.lower() == "llama-2-7b-chat-hf" or MODEL_NAME.lower() == "llama-2-13b-chat-hf":
template = LLAMA_TEMPLATE
else:
template = DEFAULT_TEMPLATE
# Use the template
prompt = template.format(job_input['prompt'])
# Streaming
streaming = job_input.get('streaming', False)
prompt = template(job_input['prompt'])
# Validate the inputs
sampling_params = job_input.get('sampling_params', None)
@@ -133,34 +129,72 @@ async def handler(job):
request_id = random_uuid()
results_generator = llm.generate(prompt, sampling_params, request_id)
# Enable HTTP Streaming
async def stream_output():
# Streaming case
async for request_output in results_generator:
prompt = request_output.prompt
text_outputs = [
prompt + output.text for output in request_output.outputs
]
ret = {"text": text_outputs}
yield ret
# Regular submission
async def submit_output():
# Non-streaming case
final_output = None
async for request_output in results_generator:
final_output = request_output
prompt = final_output.prompt
# Streaming case
async for request_output in results_generator:
prompt = request_output.prompt
text_outputs = [
prompt + output.text for output in final_output.outputs]
ret = {"outputs": text_outputs}
return ret
prompt + output.text for output in request_output.outputs
]
ret = {"text": text_outputs}
yield ret
if streaming:
return await stream_output()
async def handler(job):
'''
This is the handler function that will be called by the serverless worker.
'''
print("Job received by handler: {}".format(job))
# Get job input
job_input = job['input']
# Prompts
if MODEL_NAME.lower() == "llama-2-7b-chat-hf" or MODEL_NAME.lower() == "llama-2-13b-chat-hf":
template = LLAMA_TEMPLATE
else:
return await submit_output()
template = DEFAULT_TEMPLATE
runpod.serverless.start(
{"handler": handler, "concurrency_controller": concurrency_controller})
# Use the template
prompt = template(job_input['prompt'])
# Validate the inputs
sampling_params = job_input.get('sampling_params', None)
if sampling_params:
sampling_params = validate_sampling_params(sampling_params)
# Sampling parameters
# https://github.com/vllm-project/vllm/blob/main/vllm/sampling_params.py#L7
sampling_params = SamplingParams(**sampling_params)
else:
sampling_params = SamplingParams()
# Print the job input
print(job_input)
# Print the sampling params
print(sampling_params)
# Send request to VLLM
request_id = random_uuid()
results_generator = llm.generate(prompt, sampling_params, request_id)
# Non-streaming case
final_output = None
async for request_output in results_generator:
final_output = request_output
prompt = final_output.prompt
text_outputs = [
prompt + output.text for output in final_output.outputs]
ret = {"outputs": text_outputs}
return ret
# Start the serverless worker
if STREAMING:
print("Starting the vLLM serverless worker with streaming enabled.")
runpod.serverless.start(
{"handler": handler_streaming, "concurrency_controller": concurrency_controller})
else:
print("Starting the vLLM serverless worker with streaming disabled.")
runpod.serverless.start(
{"handler": handler, "concurrency_controller": concurrency_controller})