refactor for useage on runpod hub
This commit is contained in:
@@ -0,0 +1,45 @@
|
||||
ARG OLLAMA_VERSION=0.5.7
|
||||
|
||||
# Use an official base${OLLAMA_VERSION} image with your desired version
|
||||
FROM ollama/ollama:${OLLAMA_VERSION}
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
|
||||
# Set up the working directory
|
||||
WORKDIR /
|
||||
|
||||
RUN apt-get update --yes --quiet && DEBIAN_FRONTEND=noninteractive apt-get install --yes --quiet --no-install-recommends \
|
||||
software-properties-common \
|
||||
gpg-agent \
|
||||
build-essential apt-utils \
|
||||
&& apt-get install --reinstall ca-certificates \
|
||||
&& add-apt-repository --yes ppa:deadsnakes/ppa && apt update --yes --quiet \
|
||||
&& DEBIAN_FRONTEND=noninteractive apt-get install --yes --quiet --no-install-recommends \
|
||||
python3.11 \
|
||||
python3.11-dev \
|
||||
python3.11-distutils \
|
||||
python3.11-lib2to3 \
|
||||
python3.11-gdbm \
|
||||
python3.11-tk \
|
||||
curl \
|
||||
pip && \
|
||||
ln -s /usr/bin/python3 /usr/bin/python && \
|
||||
pip install --upgrade pip && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set the working directory
|
||||
WORKDIR /work
|
||||
|
||||
# Add my src as /work
|
||||
ADD engine.py handler.py requirements.txt start.sh test_input.json utils.py /work
|
||||
|
||||
# Set defaut ollama models directory to /runpod-volume where runpod will mount the volume by default
|
||||
ENV OLLAMA_MODELS="/runpod-volume"
|
||||
|
||||
# Install runpod and its dependencies
|
||||
RUN pip install -r requirements.txt && \
|
||||
chmod +x start.sh
|
||||
|
||||
# Set the entrypoint
|
||||
ENTRYPOINT ["/bin/sh", "-c", "/work/start.sh"]
|
||||
@@ -0,0 +1,109 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from openai import OpenAI
|
||||
from utils import JobInput
|
||||
|
||||
client = OpenAI(
|
||||
base_url='http://localhost:11434/v1/',
|
||||
|
||||
# required but ignored
|
||||
api_key='ollama',
|
||||
)
|
||||
|
||||
class OllamaEngine:
|
||||
def __init__(self):
|
||||
load_dotenv()
|
||||
print ("OllamaEngine initialized")
|
||||
|
||||
async def generate(self, job_input):
|
||||
# Get model from MODEL_NAME defauting to llama3.2:1b
|
||||
model = os.getenv("MODEL_NAME", "llama3.2:1b")
|
||||
|
||||
# Depending if prompt is a string or a list, we need to handle it differently and send it to the OpenAI API
|
||||
if isinstance(job_input.llm_input, str):
|
||||
# Buid new JobInput object with the OpenAI route and input
|
||||
openAiJob = JobInput({
|
||||
"openai_route": "/v1/completions",
|
||||
"openai_input": {
|
||||
"model": model,
|
||||
"prompt": job_input.llm_input,
|
||||
"stream": job_input.stream
|
||||
}
|
||||
})
|
||||
else:
|
||||
# Buid new JobInput object with the OpenAI route and input
|
||||
openAiJob = JobInput({
|
||||
"openai_route": "/v1/chat/completions",
|
||||
"openai_input": {
|
||||
"model": model,
|
||||
"messages": job_input.llm_input,
|
||||
"stream": job_input.stream
|
||||
}
|
||||
})
|
||||
|
||||
print ("Generating response for job_input:", job_input)
|
||||
print ("OpenAI job:", openAiJob)
|
||||
|
||||
# Create a generator that will yield the response from the OpenAI API
|
||||
openAIEngine = OllamaOpenAiEngine()
|
||||
generate = openAIEngine.generate(openAiJob)
|
||||
|
||||
# Yield the response from the OpenAI API
|
||||
async for batch in generate:
|
||||
yield batch
|
||||
|
||||
class OllamaOpenAiEngine(OllamaEngine):
|
||||
def __init__(self):
|
||||
load_dotenv()
|
||||
print ("OllamaOpenAiEngine initialized")
|
||||
|
||||
async def generate(self, job_input):
|
||||
print("Generating response for job_input:", job_input)
|
||||
|
||||
# Dump job_input to console
|
||||
openai_input = job_input.openai_input
|
||||
|
||||
# for now e just mock the response
|
||||
if job_input.openai_route == "/v1/models":
|
||||
# Async response
|
||||
async for response in self._handle_model_request():
|
||||
yield response
|
||||
elif job_input.openai_route in ["/v1/chat/completions", "/v1/completions"]:
|
||||
async for response in self._handle_chat_or_completion_request(openai_input, chat=job_input.openai_route == "/v1/chat/completions"):
|
||||
yield response
|
||||
else:
|
||||
yield {"error": "Invalid route"}
|
||||
|
||||
async def _handle_model_request(self):
|
||||
try:
|
||||
response = client.models.list()
|
||||
# build a json response from the response object
|
||||
# SyncPage[Model](data=[Model(id='llama3.2:1b', created=1737206544, object='model', owned_by='library')], object='list')\n
|
||||
yield {"object": "list", "data": [model.to_dict() for model in response.data]}
|
||||
except Exception as e:
|
||||
yield {"error": str(e)}
|
||||
|
||||
async def _handle_chat_or_completion_request(self, openai_input, chat=False):
|
||||
try:
|
||||
# Call openai.chat.completions.create or openai.completions.create based on the route
|
||||
if chat:
|
||||
response = client.chat.completions.create(**openai_input)
|
||||
else:
|
||||
response = client.completions.create(**openai_input)
|
||||
|
||||
# If streaming is False, we can just return the response
|
||||
if not openai_input.get("stream", False):
|
||||
yield response.to_dict()
|
||||
return
|
||||
|
||||
for chunk in response:
|
||||
# Log message to console
|
||||
print("Message:", chunk)
|
||||
# Return json of the chunk without any line breaks
|
||||
yield "data: " + json.dumps(chunk.to_dict(), separators=(',', ':')) + "\n\n"
|
||||
|
||||
yield "data: [DONE]"
|
||||
except Exception as e:
|
||||
yield {"error": str(e)}
|
||||
@@ -0,0 +1,32 @@
|
||||
import runpod
|
||||
from utils import JobInput
|
||||
from engine import OllamaEngine, OllamaOpenAiEngine
|
||||
|
||||
|
||||
async def handler(job: any):
|
||||
# Just dump the whole input to the console and then return an {"ok": True} response
|
||||
print('Job:', job)
|
||||
|
||||
job_input = JobInput(job["input"])
|
||||
engine_class = OllamaOpenAiEngine if job_input.openai_route else OllamaEngine
|
||||
engine = engine_class() # Instantiate the engine
|
||||
|
||||
job = engine.generate(job_input) # Call generate with job_input
|
||||
|
||||
async for batch in job:
|
||||
yield batch
|
||||
|
||||
# Original code from vllm runpod_wrapper.py
|
||||
#async def handler(job):
|
||||
# job_input = JobInput(job["input"])
|
||||
# engine = OpenAIvLLMEngine if job_input.openai_route else vllm_engine
|
||||
# results_generator = engine.generate(job_input)
|
||||
# async for batch in results_generator:
|
||||
# yield batch
|
||||
|
||||
runpod.serverless.start(
|
||||
{
|
||||
"handler": handler,
|
||||
"return_aggregate_stream": True,
|
||||
}
|
||||
)
|
||||
@@ -0,0 +1,4 @@
|
||||
runpod
|
||||
python-dotenv
|
||||
openai
|
||||
orjson==3.10.14
|
||||
@@ -0,0 +1,40 @@
|
||||
#!/bin/bash
|
||||
|
||||
cleanup() {
|
||||
echo "Cleaning up..."
|
||||
pkill -P $$ # Kill all child processes of the current script
|
||||
exit 0
|
||||
}
|
||||
|
||||
# Trap exit signals and call the cleanup function
|
||||
trap cleanup SIGINT SIGTERM
|
||||
|
||||
# Kill any existing ollama processes
|
||||
pgrep ollama | xargs kill
|
||||
|
||||
# Start the ollama server and log its output
|
||||
ollama serve 2>&1 | tee ollama.server.log &
|
||||
OLLAMA_PID=$! # Store the process ID (PID) of the background command
|
||||
|
||||
check_server_is_running() {
|
||||
echo "Checking if server is running..."
|
||||
if cat ollama.server.log | grep -q "Listening"; then
|
||||
return 0 # Success
|
||||
else
|
||||
return 1 # Failure
|
||||
fi
|
||||
}
|
||||
|
||||
# Wait for the server to start
|
||||
while ! check_server_is_running; do
|
||||
sleep 5
|
||||
done
|
||||
# IF $MODEL_NAME is set, make sure to pull the model, else just skip
|
||||
if [ -z "$MODEL_NAME" ]; then
|
||||
echo "No model name provided. Skipping model pull..."
|
||||
else
|
||||
echo "Pulling model $MODEL_NAME..."
|
||||
ollama pull $MODEL_NAME
|
||||
fi
|
||||
|
||||
python -u handler.py $1
|
||||
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"input": {
|
||||
"prompt": "How are you?"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
class JobInput:
|
||||
def __init__(self, job):
|
||||
self.llm_input = job.get("messages", job.get("prompt"))
|
||||
self.stream = job.get("stream", False)
|
||||
self.openai_route = job.get("openai_route")
|
||||
self.openai_input = job.get("openai_input")
|
||||
Reference in New Issue
Block a user