Docker-Protected HF Token, Refactor, Better Documentation
This commit is contained in:
@@ -1,3 +0,0 @@
|
||||
MODEL_NAME="mistralai/Mistral-7B-Instruct-v0.1"
|
||||
MODEL_BASE_PATH="./models"
|
||||
DISABLE_LOG_STATS=0
|
||||
+12
-9
@@ -1,6 +1,6 @@
|
||||
# Base image - Set default to CUDA 11.8
|
||||
# syntax = docker/dockerfile:1.3
|
||||
ARG WORKER_CUDA_VERSION=11.8
|
||||
FROM runpod/base:0.4.2-cuda${WORKER_CUDA_VERSION}.0 as builder
|
||||
FROM runpod/base:0.4.4-cuda${WORKER_CUDA_VERSION}.0 as builder
|
||||
|
||||
ARG WORKER_CUDA_VERSION=11.8 # Required duplicate to keep in scope
|
||||
|
||||
@@ -34,15 +34,18 @@ COPY src .
|
||||
# Setup for Option 2: Building the Image with the Model included
|
||||
ARG MODEL_NAME=""
|
||||
ARG MODEL_BASE_PATH="/runpod-volume/"
|
||||
ARG HF_TOKEN=""
|
||||
ARG QUANTIZATION=""
|
||||
RUN if [ -n "$MODEL_NAME" ]; then \
|
||||
export MODEL_BASE_PATH=$MODEL_BASE_PATH && \
|
||||
export MODEL_NAME=$MODEL_NAME && \
|
||||
python3.11 /download_model.py --model $MODEL_NAME; \
|
||||
|
||||
ENV MODEL_BASE_PATH=$MODEL_BASE_PATH \
|
||||
MODEL_NAME=$MODEL_NAME \
|
||||
QUANTIZATION=$QUANTIZATION
|
||||
|
||||
RUN --mount=type=secret,id=HF_TOKEN,required=false \
|
||||
if [ -f /run/secrets/HF_TOKEN ]; then \
|
||||
export HF_TOKEN=$(cat /run/secrets/HF_TOKEN); \
|
||||
fi && \
|
||||
if [ -n "$QUANTIZATION" ]; then \
|
||||
export QUANTIZATION=$QUANTIZATION; \
|
||||
if [ -n "$MODEL_NAME" ]; then \
|
||||
python3.11 /download_model.py --model $MODEL_NAME; \
|
||||
fi
|
||||
|
||||
# Start the handler
|
||||
|
||||
@@ -7,6 +7,23 @@
|
||||
🚀 | This serverless worker utilizes vLLM behind the scenes and is integrated into RunPod's serverless environment. It supports dynamic auto-scaling using the built-in RunPod autoscaling feature.
|
||||
</div>
|
||||
|
||||
## Table of Contents
|
||||
- [Setting up the Serverless Worker](#setting-up-the-serverless-worker)
|
||||
- [Option 1: Deploy Any Model Using Pre-Built Docker Image](#option-1-deploy-any-model-using-pre-built-docker-image)
|
||||
- [Prerequisites](#prerequisites)
|
||||
- [Environment Variables](#environment-variables)
|
||||
- [Option 2: Build Docker Image with Model Inside](#option-2-build-docker-image-with-model-inside)
|
||||
- [Arguments](#arguments)
|
||||
- [Example: Building an image with OpenChat-3.5](#example-building-an-image-with-openchat-35)
|
||||
- [(Optional) Including Huggingface Token](#optional-including-huggingface-token)
|
||||
- [Compatible Models](#compatible-models)
|
||||
- [Usage](#usage)
|
||||
- [Endpoint Model Inputs](#endpoint-model-inputs)
|
||||
- [Text Input Formats](#text-input-formats)
|
||||
- [1. `prompt`](#1-prompt)
|
||||
- [2. `messages`](#2-messages)
|
||||
- [Sampling Parameters](#sampling-parameters)
|
||||
|
||||
## Setting up the Serverless Worker
|
||||
|
||||
### Option 1: Deploy Any Model Using Pre-Built Docker Image
|
||||
@@ -21,6 +38,9 @@ Development Image: ```runpod/worker-vllm:dev```
|
||||
|
||||
</div>
|
||||
|
||||
#### Prerequisites
|
||||
- RunPod Account
|
||||
|
||||
#### Environment Variables
|
||||
|
||||
- **Required**:
|
||||
@@ -38,21 +58,42 @@ Development Image: ```runpod/worker-vllm:dev```
|
||||
- `DISABLE_LOG_REQUESTS`: Enable (`0`) or disable (`1`) request logging.
|
||||
|
||||
### Option 2: Build Docker Image with Model Inside
|
||||
To build an image with the model baked in, you must specify the following docker arguments when building the image:
|
||||
To build an image with the model baked in, you must specify the following docker arguments when building the image.
|
||||
|
||||
#### Prerequisites
|
||||
- Docker
|
||||
- Linux
|
||||
- NVIDIA GPU
|
||||
> [!NOTE]
|
||||
> We will be adding support for building on any OS without a GPU.
|
||||
|
||||
#### Arguments:
|
||||
|
||||
- **Required**
|
||||
- `MODEL_NAME`
|
||||
- **Optional**
|
||||
- `MODEL_BASE_PATH`: Defaults to `/runpod-volume` for network storage. Use `/models` or for local container storage.
|
||||
- `QUANTIZATION`
|
||||
- `HF_TOKEN`
|
||||
- `WORKER_CUDA_VERSION`: `11.8` or `12.1` (default: `11.8` due to a small amount of workers not having CUDA 12.1 support yet. `12.1` is recommended for optimal performance).
|
||||
|
||||
#### Example: Building an image with OpenChat-3.5
|
||||
```bash
|
||||
sudo docker build -t username/image:tag --build-arg MODEL_NAME="openchat/openchat_3.5" --build-arg MODEL_BASE_PATH="/models" .
|
||||
```
|
||||
|
||||
`sudo docker build -t username/image:tag --build-arg MODEL_NAME="openchat/openchat_3.5" --build-arg MODEL_BASE_PATH="/models" .`
|
||||
##### (Optional) Including Huggingface Token
|
||||
If the model you would like to deploy is private or gated, you will need to include it during build time as a Docker secret, which will protect it from being exposed in the image and on DockerHub.
|
||||
1. Enable Docker BuildKit (required for secrets).
|
||||
```bash
|
||||
export DOCKER_BUILDKIT=1
|
||||
```
|
||||
2. Export your Hugging Face token as an environment variable
|
||||
```bash
|
||||
export HF_TOKEN="your_secret_value_here"
|
||||
```
|
||||
2. Add the token as a secret when building
|
||||
```bash
|
||||
docker build -t username/image:tag --secret id=HF_TOKEN --build-arg MODEL_NAME="openchat/openchat_3.5" .
|
||||
```
|
||||
|
||||
### Compatible Models
|
||||
|
||||
@@ -81,7 +122,8 @@ And any other models supported by vLLM 0.2.6.
|
||||
Ensure that you have Docker installed and properly set up before running the docker build commands. Once built, you can deploy this serverless worker in your desired environment with confidence that it will automatically scale based on demand. For further inquiries or assistance, feel free to contact our support team.
|
||||
|
||||
|
||||
## Model Inputs
|
||||
## Usage
|
||||
### Endpoint Model Inputs
|
||||
You may either use a `prompt` or a list of `messages` as input. If you use `messages`, the model's chat template will be applied to the messages automatically, so the model must have one. If you use `prompt`, you may optionally apply the model's chat template to the prompt by setting `apply_chat_template` to `true`.
|
||||
| Argument | Type | Default | Description |
|
||||
|-----------------------|----------------------|--------------------|--------------------------------------------------------------------------------------------------------|
|
||||
@@ -92,17 +134,26 @@ You may either use a `prompt` or a list of `messages` as input. If you use `mess
|
||||
| `stream` | bool | False | Whether to enable streaming of output. If True, responses are streamed as they are generated. |
|
||||
| `batch_size` | int | DEFAULT_BATCH_SIZE | The number of tokens to stream every HTTP POST call. |
|
||||
|
||||
### Messages Format
|
||||
### Text Input Formats
|
||||
You may either use a `prompt` or a list of `messages` as input.
|
||||
#### 1. `prompt`
|
||||
The prompt string can be any string, and the model's chat template will not be applied to it unless `apply_chat_template` is set to `true`, in which case it will be treated as a user message.
|
||||
|
||||
Example:
|
||||
```json
|
||||
"prompt": "..."
|
||||
```
|
||||
#### 2. `messages`
|
||||
Your list can contain any number of messages, and each message can have any role from the following list:
|
||||
- `user`
|
||||
- `assistant`
|
||||
- `system`
|
||||
|
||||
The model's chat template will be applied to the messages automatically.
|
||||
The model's chat template will be applied to the messages automatically, so the model must have one.
|
||||
|
||||
Example:
|
||||
```json
|
||||
[
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "..."
|
||||
@@ -115,7 +166,7 @@ Example:
|
||||
"role": "assistant",
|
||||
"content": "..."
|
||||
}
|
||||
]
|
||||
]
|
||||
```
|
||||
|
||||
### Sampling Parameters
|
||||
|
||||
+6
-4
@@ -1,20 +1,22 @@
|
||||
from typing import Union
|
||||
|
||||
DEFAULT_BATCH_SIZE = 30
|
||||
DEFAULT_MAX_CONCURRENCY = 300
|
||||
|
||||
sampling_param_types = {
|
||||
SAMPLING_PARAM_TYPES = {
|
||||
"n": int,
|
||||
"best_of": int,
|
||||
"presence_penalty": float,
|
||||
"frequency_penalty": float,
|
||||
"repetition_penalty": float,
|
||||
"temperature": float,
|
||||
"temperature": Union[float, int],
|
||||
"top_p": float,
|
||||
"top_k": int,
|
||||
"min_p": float,
|
||||
"use_beam_search": bool,
|
||||
"length_penalty": float,
|
||||
"early_stopping": (bool, str),
|
||||
"stop": (str, list),
|
||||
"early_stopping": Union[bool, str],
|
||||
"stop": Union[str, list],
|
||||
"stop_token_ids": list,
|
||||
"ignore_eos": bool,
|
||||
"max_tokens": int,
|
||||
|
||||
@@ -79,11 +79,3 @@ class vLLMEngine:
|
||||
quantization = os.getenv("QUANTIZATION", "").lower()
|
||||
return quantization if quantization in ["awq", "squeezellm", "gptq"] else None
|
||||
|
||||
def concurrency_modifier(self, current_concurrency):
|
||||
n_current_jobs = self._get_n_current_jobs()
|
||||
requested_concurrency = max(0, self.serverless_config.max_concurrency - n_current_jobs)
|
||||
if not self.config["disable_log_stats"]:
|
||||
logging.info("Current Jobs: %s", n_current_jobs)
|
||||
logging.info("Concurrency Modifier Requested Jobs: %s", requested_concurrency)
|
||||
return requested_concurrency
|
||||
|
||||
|
||||
+16
-23
@@ -1,31 +1,29 @@
|
||||
#!/usr/bin/env python
|
||||
from typing import Generator
|
||||
from vllm.utils import random_uuid
|
||||
import runpod
|
||||
from utils import validate_sampling_params, random_uuid
|
||||
from utils import validate_sampling_params
|
||||
from engine import vLLMEngine
|
||||
|
||||
vllm_engine = vLLMEngine()
|
||||
|
||||
async def handler(job: dict) -> Generator[dict, None, None]:
|
||||
job_input = job["input"]
|
||||
llm_input = job_input.get("messages", job_input.get("prompt"))
|
||||
apply_chat_template = job_input.get("apply_chat_template", False)
|
||||
|
||||
if apply_chat_template or isinstance(llm_input, list):
|
||||
if job_input.get("apply_chat_template", False) or isinstance(llm_input, list):
|
||||
llm_input = vllm_engine.tokenizer.apply_chat_template(llm_input)
|
||||
|
||||
stream = job_input.get("stream", False)
|
||||
batch_size = job_input.get("batch_size", vllm_engine.serverless_config.default_batch_size)
|
||||
sampling_params = job_input.get("sampling_params", {})
|
||||
|
||||
validated_params = validate_sampling_params(sampling_params)
|
||||
batch_size = job_input.get("batch_size", vllm_engine.serverless_config.batch_size)
|
||||
validated_params = validate_sampling_params(job_input.get("sampling_params", {}))
|
||||
request_id = random_uuid()
|
||||
|
||||
results_generator = vllm_engine.llm.generate(
|
||||
llm_input, validated_params, request_id
|
||||
)
|
||||
|
||||
batch = {"tokens": []}
|
||||
last_output_text = ""
|
||||
n_input_tokens, is_first_output = 0, True
|
||||
batch, last_output_text, n_input_tokens, is_first_output = {"tokens": []}, "", 0, True
|
||||
|
||||
|
||||
async for request_output in results_generator:
|
||||
if is_first_output: # Count input tokens only once
|
||||
@@ -34,29 +32,24 @@ async def handler(job: dict) -> Generator[dict, None, None]:
|
||||
|
||||
for output in request_output.outputs:
|
||||
if stream:
|
||||
batch["tokens"].append(output.text[len(last_output_text):])
|
||||
|
||||
batch["tokens"].append(
|
||||
output.text[len(last_output_text):]
|
||||
)
|
||||
finished = request_output.finished
|
||||
if len(batch["tokens"]) >= batch_size or finished:
|
||||
if len(batch["tokens"]) >= batch_size:
|
||||
batch["usage"] = {
|
||||
"input": n_input_tokens,
|
||||
"output": len(output.token_ids),
|
||||
}
|
||||
batch["finished"] = finished
|
||||
yield batch
|
||||
batch = {"tokens": []}
|
||||
|
||||
last_output_text = output.text
|
||||
|
||||
if not stream:
|
||||
yield {"tokens": [last_output_text],
|
||||
"usage": {
|
||||
"input": n_input_tokens,
|
||||
"output": len(output.token_ids),
|
||||
},
|
||||
"finished": True}
|
||||
batch["tokens"].append(last_output_text)
|
||||
|
||||
if len(batch["tokens"]) > 0:
|
||||
batch["usage"] = {"input": n_input_tokens, "output": len(output.token_ids)}
|
||||
yield batch
|
||||
|
||||
runpod.serverless.start(
|
||||
{
|
||||
|
||||
+10
-34
@@ -2,50 +2,26 @@ import os
|
||||
import logging
|
||||
from typing import Any, Dict
|
||||
from vllm import SamplingParams
|
||||
from vllm.utils import random_uuid
|
||||
from constants import sampling_param_types, DEFAULT_BATCH_SIZE, DEFAULT_MAX_CONCURRENCY
|
||||
from constants import SAMPLING_PARAM_TYPES, DEFAULT_BATCH_SIZE, DEFAULT_MAX_CONCURRENCY
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
|
||||
class ServerlessConfig:
|
||||
def __init__(self):
|
||||
self._max_concurrency = int(
|
||||
os.environ.get("MAX_CONCURRENCY", DEFAULT_MAX_CONCURRENCY)
|
||||
)
|
||||
self._default_batch_size = int(
|
||||
os.environ.get("DEFAULT_BATCH_SIZE", DEFAULT_BATCH_SIZE)
|
||||
)
|
||||
|
||||
@property
|
||||
def max_concurrency(self):
|
||||
return self._max_concurrency
|
||||
|
||||
@property
|
||||
def default_batch_size(self):
|
||||
return self._default_batch_size
|
||||
|
||||
self.max_concurrency = int(os.getenv("MAX_CONCURRENCY", DEFAULT_MAX_CONCURRENCY))
|
||||
self.batch_size = int(os.getenv("BATCH_SIZE", DEFAULT_BATCH_SIZE))
|
||||
|
||||
def validate_sampling_params(params: Dict[str, Any]) -> SamplingParams:
|
||||
validated_params = {}
|
||||
|
||||
invalid_params = []
|
||||
for key, value in params.items():
|
||||
expected_type = sampling_param_types.get(key)
|
||||
if value is None:
|
||||
validated_params[key] = None
|
||||
continue
|
||||
|
||||
if expected_type is None:
|
||||
continue
|
||||
|
||||
if isinstance(expected_type, tuple):
|
||||
casted_value = next(
|
||||
(t(value) for t in expected_type if isinstance(value, t)), None
|
||||
)
|
||||
expected_type = SAMPLING_PARAM_TYPES.get(key)
|
||||
if expected_type and isinstance(value, expected_type):
|
||||
validated_params[key] = value
|
||||
else:
|
||||
casted_value = value if isinstance(value, expected_type) else None
|
||||
invalid_params.append(key)
|
||||
|
||||
if casted_value is not None:
|
||||
validated_params[key] = casted_value
|
||||
if len(invalid_params) > 0:
|
||||
logging.warning("Ignoring invalid sampling params: %s", invalid_params)
|
||||
|
||||
return SamplingParams(**validated_params)
|
||||
Reference in New Issue
Block a user