From 401a4637faf9b9d21cd6ebe4bf7d99a1006c0eb6 Mon Sep 17 00:00:00 2001 From: Jorg Doku Date: Thu, 3 Aug 2023 10:05:51 -0500 Subject: [PATCH] latest version --- Dockerfile | 23 +++++--- builder/requirements.txt | 3 +- src/download_model.py | 21 +++++--- src/handler.py | 114 +++++++++++++++++++++++++-------------- 4 files changed, 106 insertions(+), 55 deletions(-) diff --git a/Dockerfile b/Dockerfile index d82bd15..7bdf32d 100644 --- a/Dockerfile +++ b/Dockerfile @@ -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 diff --git a/builder/requirements.txt b/builder/requirements.txt index 0b2f19e..c3a3ad9 100644 --- a/builder/requirements.txt +++ b/builder/requirements.txt @@ -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 diff --git a/src/download_model.py b/src/download_model.py index 47fcbe4..9eb5299 100644 --- a/src/download_model.py +++ b/src/download_model.py @@ -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 +) diff --git a/src/handler.py b/src/handler.py index 578e0a2..142bb2d 100644 --- a/src/handler.py +++ b/src/handler.py @@ -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})