latest version
This commit is contained in:
+16
-7
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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})
|
||||
|
||||
Reference in New Issue
Block a user