diff --git a/Dockerfile b/Dockerfile index 9ea8152..87e92e7 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,6 +1,8 @@ # Base image -# The following docker base image is recommended by VLLM: nvcr.io/nvidia/pytorch:22.12-py3 -FROM runpod/pytorch:3.10-2.0.0-117 +# The following docker base image is recommended by VLLM: +# FROM runpod/pytorch:2.0.1-py3.10-cuda11.8.0-devel +# FROM nvcr.io/nvidia/pytorch:22.12-py3 +FROM runpod/pytorch:2.0.1-py3.10-cuda11.8.0-devel # Use bash shell with pipefail option SHELL ["/bin/bash", "-o", "pipefail", "-c"] @@ -9,6 +11,9 @@ SHELL ["/bin/bash", "-o", "pipefail", "-c"] WORKDIR / # Update and upgrade the system packages (Worker Template) +ARG DEBIAN_FRONTEND=noninteractive +RUN pip uninstall torch -y +RUN pip install torch==2.0.1 -f https://download.pytorch.org/whl/cu118 COPY builder/setup.sh /setup.sh RUN chmod +x /setup.sh && \ /setup.sh && \ diff --git a/builder/requirements.txt b/builder/requirements.txt index 4eae5c2..1356e8b 100644 --- a/builder/requirements.txt +++ b/builder/requirements.txt @@ -3,7 +3,7 @@ # 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_beta#egg=runpod -vllm @ git+https://github.com/runpod/vllm@runpod#egg=vllm -torch==2.0.1 +vllm==0.1.2 +runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod diff --git a/src/handler.py b/src/handler.py index 32fd65e..0d41d4a 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,121 +1,102 @@ #!/usr/bin/env python ''' Contains the handler function that will be called by the serverless. ''' from typing import Dict - -import json -import requests -from urllib.parse import urljoin - import runpod -from vllm.entrypoints.runpod.api_server import start_vllm_runpod # Start the VLLM serving layer on our RunPod worker. -vllm = start_vllm_runpod( - served_model='facebook/opt-125m', port=443, host='127.0.0.1') +from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs +from vllm.utils import random_uuid +# Prepare the model and tokenizer +MODEL = 'lmsys/vicuna-13b-v1.3' +TOKENIZER = 'hf-internal-testing/llama-tokenizer' -def prepare_request(event: Dict) -> Dict: - """ - # Pre-processing Steps for the Prompt - # Include any necessary code here to pre-process the prompt. +# Prepare the engine's arguments +engine_args = AsyncEngineArgs( + model=MODEL, + tokenizer=TOKENIZER, + tokenizer_mode= "auto", + tensor_parallel_size= 1, + dtype = "auto", + seed = 0, + worker_use_ray=False, +) +llm = AsyncLLMEngine.from_engine_args(engine_args) - # Example: - # Step 1: Clean the prompt - # cleaned_prompt = clean_text(prompt) +# Validation +def validate_sampling_params(sampling_params): + def validate_int(value, default): + try: + return int(value) + except (TypeError, ValueError): + return default - # Step 2: Apply specific transformations - # transformed_prompt = apply_transformations(tokenized_prompt) + def validate_float(value, default): + try: + return float(value) + except (TypeError, ValueError): + return default - # Step 3: Format the prompt for the model - # formatted_prompt = format_for_model(transformed_prompt) - """ + def validate_bool(value, default): + if isinstance(value, bool): + return value + return default - request_data = { - 'url': urljoin("http://127.0.0.1:443/", event['llm_endpoint']['url']), - 'headers': { - "Content-Type": "application/json" - }, - 'request_type': event['llm_endpoint']['request_type'], - 'body_data': event['llm_body'] + n = validate_int(sampling_params.get('n'), 1) + best_of = validate_int(sampling_params.get('best_of'), None) + presence_penalty = validate_float(sampling_params.get('presence_penalty'), 0.0) + frequency_penalty = validate_float(sampling_params.get('frequency_penalty'), 0.0) + temperature = validate_float(sampling_params.get('temperature'), 1.0) + 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') + 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) + + return { + 'n': n, + 'best_of': best_of, + 'presence_penalty': presence_penalty, + 'frequency_penalty': frequency_penalty, + 'temperature': temperature, + 'top_p': top_p, + 'top_k': top_k, + 'use_beam_search': use_beam_search, + 'stop': stop, + 'ignore_eos': ignore_eos, + 'max_tokens': max_tokens, + 'logprobs': logprobs, } - return request_data - - -def make_vllm_request(request_data: Dict) -> requests.Response: - url = request_data['url'] - headers = request_data['headers'] - request_type = request_data['request_type'] - - if request_type.lower() == 'post': - body_data = request_data['body_data'] - response = requests.post(url, headers=headers, - data=json.dumps(body_data)) - elif request_type.lower() == 'get': - response = requests.get(url, headers=headers) - else: - raise ValueError(f"Invalid request type: {request_type}") - - return response - - -def process_response(response: requests.Response) -> Dict: - # Process the json response. - response_data = response.json() - - """ - # Additional Post-processing Steps - # You can include any necessary code here to process the LLM's generated output. - - # Example: - # Step 1: Extract relevant information - # result = llm_output['data']['result'] - # relevant_info = result['info'] - - # Step 2: Clean the data - # cleaned_data = preprocess(relevant_info) - - # Step 3: Apply transformations or filters - # transformed_data = apply_transformations(cleaned_data) - - # Step 4: Finalize the output - # final_output = format_output(transformed_data) - - # Return the final processed output - return final_output - """ - - return response_data - - -def handler(event): +async def handler(job): ''' This is the handler function that will be called by the serverless worker. ''' - # Prepare the request for vllm. - request_data = prepare_request(event) + # Prompts + prompts = job['prompts'] - # Make the request. - response = make_vllm_request(request_data) + # Validate the inputs + sampling_params = job['sampling_params'] + sampling_params = validate_sampling_params(sampling_params) - # Process the response from vllm. - response_data = process_response(response) + # Sampling parameters + # https://github.com/vllm-project/vllm/blob/main/vllm/sampling_params.py#L7 + sampling_params = SamplingParams(**sampling_params) - return response_data + # Send request to VLLM + request_id = random_uuid() + results_generator = llm.generate(prompts, sampling_params, request_id) + # Non-streaming case + final_output = None + async for request_output in results_generator: + final_output = request_output -""" -Provide access to our custom handler, which allows us to incorporate pre-processing and -post-processing steps into the prompt. This custom handler enhances the functionality -of our program by allowing us to perform additional tasks before and after the prompt execution. + prompt = final_output.prompt + text_outputs = [prompt + output.text for output in final_output.outputs] + ret = {"text": text_outputs} + return ret -Furthermore, we pass the 'runpod vllm' instance to ensure efficient auto-scaling based on -the usage of the vllm (very large language model). This inclusion enables the program to -dynamically adjust its resource allocation to accommodate the demands of the vllm, -optimizing its performance and scalability. -""" -runpod.serverless.start( - {"handler": handler, "vllm": vllm}, serverless_llm=True) - -# Start the VLLM server -vllm.start() +runpod.serverless.start({"handler": handler})