This commit is contained in:
Jorg Doku
2023-07-20 18:45:08 -05:00
parent ff59fb6b8f
commit 9fe114c049
5 changed files with 49 additions and 10 deletions
+3
View File
@@ -0,0 +1,3 @@
{
"python.linting.enabled": true
}
+12 -1
View File
@@ -30,6 +30,17 @@ RUN --mount=type=cache,target=/root/.cache/pip \
ADD src . ADD src .
# Quick temporary updates # Quick temporary updates
RUN pip install git+https://github.com/runpod/runpod-python@multijob2#egg=runpod --compile RUN pip install git+https://github.com/runpod/runpod-python@multijob2#egg=runpod --compile
# Prepare the models inside the docker image
ARG HUGGING_FACE_HUB_TOKEN=NONE
ENV HUGGING_FACE_HUB_TOKEN=$HUGGING_FACE_HUB_TOKEN
ENV DOWNLOAD_7B_MODEL=YES
# ENV DOWNLOAD_13B_MODEL=1
# Download the models
RUN mkdir -p /model
RUN DOWNLOAD_7B_MODEL=$DOWNLOAD_7B_MODEL HUGGING_FACE_HUB_TOKEN=$HUGGING_FACE_HUB_TOKEN python -u /download_model.py
# Start the handler
CMD python -u /handler.py CMD python -u /handler.py
+2 -1
View File
@@ -6,4 +6,5 @@
# vllm @ git+https://github.com/vllm-project/vllm.git@2b7d3aca2e1dd25fe26424f57c051af3b823cd71 # vllm @ git+https://github.com/vllm-project/vllm.git@2b7d3aca2e1dd25fe26424f57c051af3b823cd71
# runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod # runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod
vllm==0.1.2 vllm==0.1.2
runpod @ git+https://github.com/runpod/runpod-python@multijob#egg=runpod huggingface-hub==0.16.4
runpod @ git+https://github.com/runpod/runpod-python@multijob2#egg=runpod
+23
View File
@@ -0,0 +1,23 @@
import os
from huggingface_hub import snapshot_download
# Get the hugging face token
HUGGING_FACE_HUB_TOKEN = os.environ['HUGGING_FACE_HUB_TOKEN']
DOWNLOAD_7B_MODEL = os.environ.get('DOWNLOAD_7B_MODEL', None)
DOWNLOAD_13B_MODEL = os.environ.get('DOWNLOAD_13B_MODEL', None)
# Download the 7B
if HUGGING_FACE_HUB_TOKEN and DOWNLOAD_7B_MODEL:
snapshot_download(
"meta-llama/Llama-2-7b-chat-hf",
local_dir="/model/Llama-2-7b-chat-hf",
token=HUGGING_FACE_HUB_TOKEN
)
# Download the 13B
if HUGGING_FACE_HUB_TOKEN and DOWNLOAD_13B_MODEL:
snapshot_download(
"meta-llama/Llama-2-13b-chat-hf",
local_dir="/model/Llama-2-13b-chat-hf",
token=HUGGING_FACE_HUB_TOKEN
)
+9 -8
View File
@@ -1,23 +1,19 @@
#!/usr/bin/env python #!/usr/bin/env python
''' Contains the handler function that will be called by the serverless. ''' ''' Contains the handler function that will be called by the serverless. '''
import json
import types
from typing import AsyncGenerator, Dict
# Start the VLLM serving layer on our RunPod worker. # Start the VLLM serving layer on our RunPod worker.
from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs
from vllm.utils import random_uuid from vllm.utils import random_uuid
import runpod import runpod
import asyncio
# Prepare the model and tokenizer # Prepare the model and tokenizer
MODEL = 'facebook/opt-125m' MODEL = '/model/Llama-2-7b-chat-hf'
# TOKENIZER = 'hf-internal-testing/llama-tokenizer' TOKENIZER = 'hf-internal-testing/llama-tokenizer'
# Prepare the engine's arguments # Prepare the engine's arguments
engine_args = AsyncEngineArgs( engine_args = AsyncEngineArgs(
model=MODEL, model=MODEL,
#tokenizer=TOKENIZER, tokenizer=TOKENIZER,
tokenizer_mode= "auto", tokenizer_mode= "auto",
tensor_parallel_size= 1, tensor_parallel_size= 1,
dtype = "auto", dtype = "auto",
@@ -28,7 +24,6 @@ engine_args = AsyncEngineArgs(
# Create the vLLM asynchronous engine # Create the vLLM asynchronous engine
llm = AsyncLLMEngine.from_engine_args(engine_args) llm = AsyncLLMEngine.from_engine_args(engine_args)
def handler_fully_utilized() -> bool: def handler_fully_utilized() -> bool:
# Compute pending sequences # Compute pending sequences
total_pending_sequences = len(llm.engine.scheduler.waiting) + len(llm.engine.scheduler.swapped) total_pending_sequences = len(llm.engine.scheduler.waiting) + len(llm.engine.scheduler.swapped)
@@ -109,6 +104,12 @@ async def handler(job):
else: else:
sampling_params = SamplingParams() sampling_params = SamplingParams()
# Print the job input
print(job_input)
# Print the sampling params
print(sampling_params)
# Send request to VLLM # Send request to VLLM
request_id = random_uuid() request_id = random_uuid()
results_generator = llm.generate(prompt, sampling_params, request_id) results_generator = llm.generate(prompt, sampling_params, request_id)