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 .
# 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
+2 -1
View File
@@ -6,4 +6,5 @@
# vllm @ git+https://github.com/vllm-project/vllm.git@2b7d3aca2e1dd25fe26424f57c051af3b823cd71
# runpod @ git+https://github.com/runpod/runpod-python@vllm#egg=runpod
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
''' 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.
from vllm import AsyncLLMEngine, SamplingParams, AsyncEngineArgs
from vllm.utils import random_uuid
import runpod
import asyncio
# Prepare the model and tokenizer
MODEL = 'facebook/opt-125m'
# TOKENIZER = 'hf-internal-testing/llama-tokenizer'
MODEL = '/model/Llama-2-7b-chat-hf'
TOKENIZER = 'hf-internal-testing/llama-tokenizer'
# Prepare the engine's arguments
engine_args = AsyncEngineArgs(
model=MODEL,
#tokenizer=TOKENIZER,
tokenizer=TOKENIZER,
tokenizer_mode= "auto",
tensor_parallel_size= 1,
dtype = "auto",
@@ -28,7 +24,6 @@ engine_args = AsyncEngineArgs(
# Create the vLLM asynchronous engine
llm = AsyncLLMEngine.from_engine_args(engine_args)
def handler_fully_utilized() -> bool:
# Compute pending sequences
total_pending_sequences = len(llm.engine.scheduler.waiting) + len(llm.engine.scheduler.swapped)
@@ -109,6 +104,12 @@ async def handler(job):
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)