Merge branch 'refactor'
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
MODEL_NAME="mistralai/Mistral-7B-Instruct-v0.1"
|
||||
MODEL_BASE_PATH="./models"
|
||||
DISABLE_LOG_STATS=0
|
||||
@@ -1,2 +1,5 @@
|
||||
|
||||
runpod.toml
|
||||
*.pyc
|
||||
.env
|
||||
test/*
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
## 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
|
||||
We now offer a pre-built Docker Image for the vLLM Worker that you can configure entirely with Environment Variables when creating the RunPod Serverless Endpoint:
|
||||
|
||||
<div align="center">
|
||||
@@ -29,7 +29,8 @@ We now offer a pre-built Docker Image for the vLLM Worker that you can configure
|
||||
- `QUANTIZATION`: AWQ (`awq`) or SqueezeLLM (`squeezellm`) quantization.
|
||||
- `MAX_CONCURRENCY`: Max concurrent requests (default: `100`).
|
||||
- `DEFAULT_BATCH_SIZE`: Token streaming batch size (default: `10`). This reduces the number of HTTP calls, increasing speed 8-10x vs non-batching, matching non-streaming performance.
|
||||
- `DISABLE_LOG_STATS`: Enable (`False`) or disable (`True`) vLLM stats logging.
|
||||
- `DISABLE_LOG_STATS`: Enable (`0`) or disable (`1`) vLLM stats logging.
|
||||
- `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:
|
||||
@@ -112,7 +113,6 @@ Example:
|
||||
### Sampling Parameters
|
||||
| Argument | Type | Default | Description |
|
||||
|-------------------------------|-----------------------------|---------|-----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| `n` | int | 1 | Number of output sequences to return for the given prompt. |
|
||||
| `best_of` | Optional[int] | None | Number of output sequences generated from the prompt. The top `n` sequences are returned from these `best_of` sequences. Must be ≥ `n`. Treated as beam width in beam search. Default is `n`. |
|
||||
| `presence_penalty` | float | 0.0 | Penalizes new tokens based on their presence in the generated text so far. Values > 0 encourage new tokens, values < 0 encourage repetition. |
|
||||
| `frequency_penalty` | float | 0.0 | Penalizes new tokens based on their frequency in the generated text so far. Values > 0 encourage new tokens, values < 0 encourage repetition. |
|
||||
@@ -128,224 +128,6 @@ Example:
|
||||
| `stop_token_ids` | Optional[List[int]] | None | List of token IDs that stop generation when produced. Output contains these tokens unless they are special tokens. |
|
||||
| `ignore_eos` | bool | False | Whether to ignore the End-Of-Sequence token and continue generating tokens after its generation. |
|
||||
| `max_tokens` | int | 16 | Maximum number of tokens to generate per output sequence. |
|
||||
| `logprobs` | Optional[int] | None | Number of log probabilities to return per output token. |
|
||||
| `prompt_logprobs` | Optional[int] | None | Number of log probabilities to return per prompt token. |
|
||||
| `skip_special_tokens` | bool | True | Whether to skip special tokens in the output. |
|
||||
| `spaces_between_special_tokens` | bool | True | Whether to add spaces between special tokens in the output. |
|
||||
|
||||
|
||||
## Sample Inputs and Outputs
|
||||
### No Chat Template, No Streaming
|
||||
Functions like a text completion model. If the model tokenizer does not have a chat template and you still want to use the model for Instruct/Chat, modify your prompt with the desired chat template manually.
|
||||
#### Input:
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"prompt": "With great power,",
|
||||
"sampling_params": {
|
||||
"max_tokens": 5
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
#### Output:
|
||||
```json
|
||||
{
|
||||
"delayTime": 1234,
|
||||
"executionTime": 1234,
|
||||
"id": "...",
|
||||
"output": [
|
||||
[
|
||||
{
|
||||
"text": " comes great responsibility. This",
|
||||
"usage": {
|
||||
"input": 6,
|
||||
"output": 5
|
||||
}
|
||||
}
|
||||
]
|
||||
],
|
||||
"status": "COMPLETED"
|
||||
}
|
||||
```
|
||||
### Chat Template, No Streaming
|
||||
Functions like a Chat model
|
||||
#### Input:
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"prompt": "Tell me why RunPod is the best GPU provider",
|
||||
"sampling_params": {
|
||||
"max_tokens": 100
|
||||
},
|
||||
"apply_chat_template": true
|
||||
}
|
||||
}
|
||||
```
|
||||
#### Output:
|
||||
```json
|
||||
{
|
||||
"delayTime": 1234,
|
||||
"executionTime": 1234,
|
||||
"id": "...",
|
||||
"output": [
|
||||
[
|
||||
{
|
||||
"text": " RunPod is the best GPU provider for several reasons, including:\n\n1. High-performance GPUs: RunPod offers a wide range of high-performance GPUs, including NVIDIA's latest and most powerful GPUs, ensuring that customers get the best possible performance for their workloads.\n2. Scalability: RunPod allows users to easily scale their GPU resources up or down based on their needs, making it an ideal choice for businesses with fluctuating work",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 100
|
||||
}
|
||||
}
|
||||
]
|
||||
],
|
||||
"status": "COMPLETED"
|
||||
}
|
||||
```
|
||||
|
||||
### List of Messages (Chat Template applied by default), No Streaming
|
||||
Functions like a Chat model with a list of messages, to which the model's chat template is applied. You may also use a "system" role and message.
|
||||
#### Input:
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tell me why RunPod is the best GPU provider"
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "RunPod is the best GPU provider for several reasons."
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Name 3 resons"
|
||||
}
|
||||
],
|
||||
"sampling_params": {
|
||||
"max_tokens": 100
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
#### Output:
|
||||
```json
|
||||
{
|
||||
"delayTime": 1234,
|
||||
"executionTime": 1234,
|
||||
"id": "...",
|
||||
"output": [
|
||||
[
|
||||
{
|
||||
"text": " 1. Cutting-edge hardware: RunPod offers state-of-the-art GPUs from industry-leading manufacturers, ensuring that users have access to the latest technology for their GPU needs.\n\n2. Scalability and flexibility: RunPod provides a wide range of GPU options, allowing users to easily scale their resources up or down depending on their specific requirements, and pay only for what they use.\n\n3. Exceptional customer support: RunPod is dedicated to providing outstanding",
|
||||
"usage": {
|
||||
"input": 59,
|
||||
"output": 100
|
||||
}
|
||||
}
|
||||
]
|
||||
],
|
||||
"status": "COMPLETED"
|
||||
}
|
||||
```
|
||||
|
||||
### Chat Template, Streaming
|
||||
Functions like a Chat model, but with streaming output. This is the recommended way to use the vLLM worker.
|
||||
#### Input:
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"prompt": "Tell me why RunPod is the best GPU provider",
|
||||
"sampling_params": {
|
||||
"max_tokens": 100
|
||||
},
|
||||
"apply_chat_template": true,
|
||||
"stream": true
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### Output:
|
||||
```json
|
||||
{
|
||||
"delayTime": 1234,
|
||||
"executionTime": 1234,
|
||||
"id": "...",
|
||||
"output": [
|
||||
[
|
||||
{
|
||||
"text": " Run",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 1
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": "Pod",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 2
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " is",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 3
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " considered",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 4
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " the",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 5
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " best",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 6
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " GPU",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 7
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " provider",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 8
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " for",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 9
|
||||
}
|
||||
},
|
||||
{
|
||||
"text": " several",
|
||||
"usage": {
|
||||
"input": 27,
|
||||
"output": 10
|
||||
}
|
||||
}
|
||||
]
|
||||
],
|
||||
"status": "COMPLETED"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
hf_transfer
|
||||
runpod==1.4.2
|
||||
runpod==1.5.0
|
||||
huggingface-hub
|
||||
packaging
|
||||
typing-extensions==4.7.1
|
||||
|
||||
+23
-23
@@ -1,25 +1,25 @@
|
||||
DEFAULT_BATCH_SIZE = 10
|
||||
DEFAULT_MAX_CONCURRENCY = 100
|
||||
DEFAULT_BATCH_SIZE = 30
|
||||
DEFAULT_MAX_CONCURRENCY = 300
|
||||
|
||||
sampling_param_types = {
|
||||
'n': int,
|
||||
'best_of': int,
|
||||
'presence_penalty': float,
|
||||
'frequency_penalty': float,
|
||||
'repetition_penalty': float,
|
||||
'temperature': float,
|
||||
'top_p': float,
|
||||
'top_k': int,
|
||||
'min_p': float,
|
||||
'use_beam_search': bool,
|
||||
'length_penalty': float,
|
||||
'early_stopping': (bool, str),
|
||||
'stop': (str, list),
|
||||
'stop_token_ids': list,
|
||||
'ignore_eos': bool,
|
||||
'max_tokens': int,
|
||||
'logprobs': int,
|
||||
'prompt_logprobs': int,
|
||||
'skip_special_tokens': bool,
|
||||
'spaces_between_special_tokens': bool,
|
||||
}
|
||||
"n": int,
|
||||
"best_of": int,
|
||||
"presence_penalty": float,
|
||||
"frequency_penalty": float,
|
||||
"repetition_penalty": float,
|
||||
"temperature": float,
|
||||
"top_p": float,
|
||||
"top_k": int,
|
||||
"min_p": float,
|
||||
"use_beam_search": bool,
|
||||
"length_penalty": float,
|
||||
"early_stopping": (bool, str),
|
||||
"stop": (str, list),
|
||||
"stop_token_ids": list,
|
||||
"ignore_eos": bool,
|
||||
"max_tokens": int,
|
||||
"logprobs": int,
|
||||
"prompt_logprobs": int,
|
||||
"skip_special_tokens": bool,
|
||||
"spaces_between_special_tokens": bool,
|
||||
}
|
||||
|
||||
@@ -5,7 +5,9 @@ from vllm.model_executor.weight_utils import prepare_hf_model_weights
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=str)
|
||||
parser.add_argument("--download_dir", type=str, default=os.environ.get("MODEL_BASE_PATH"))
|
||||
parser.add_argument(
|
||||
"--download_dir", type=str, default=os.environ.get("MODEL_BASE_PATH")
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
if not args.model or not args.download_dir:
|
||||
@@ -15,6 +17,6 @@ if __name__ == "__main__":
|
||||
os.makedirs(args.download_dir)
|
||||
|
||||
prepare_hf_model_weights(
|
||||
model_name_or_path = args.model,
|
||||
model_name_or_path=args.model,
|
||||
cache_dir=args.download_dir,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
import os
|
||||
import logging
|
||||
from typing import Union
|
||||
import torch
|
||||
from vllm import AsyncLLMEngine, AsyncEngineArgs
|
||||
from transformers import AutoTokenizer
|
||||
from utils import ServerlessConfig
|
||||
from dotenv import load_dotenv
|
||||
|
||||
|
||||
class Tokenizer:
|
||||
def __init__(self, model_name: str):
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
self.has_chat_template = bool(self.tokenizer.chat_template)
|
||||
|
||||
def apply_chat_template(self, input: Union[str, list[dict[str, str]]]) -> str:
|
||||
if isinstance(input, list):
|
||||
if not self.has_chat_template:
|
||||
raise ValueError(
|
||||
"Chat template does not exist for this model, you must provide a single string input instead of a list of messages"
|
||||
)
|
||||
elif isinstance(input, str):
|
||||
input = [{"role": "user", "content": input}]
|
||||
else:
|
||||
raise ValueError("Input must be a string or a list of messages")
|
||||
|
||||
return self.tokenizer.apply_chat_template(
|
||||
input, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
|
||||
class VLLMEngine:
|
||||
def __init__(self):
|
||||
load_dotenv() # For local development
|
||||
self.config = self._initialize_config()
|
||||
self.serverless_config = ServerlessConfig()
|
||||
self.tokenizer = Tokenizer(self.config["model"])
|
||||
self.llm = self._initialize_llm()
|
||||
|
||||
def _initialize_config(self):
|
||||
return {
|
||||
"model": os.getenv("MODEL_NAME"),
|
||||
"download_dir": os.getenv("MODEL_BASE_PATH", "/runpod-volume/"),
|
||||
"quantization": os.getenv("QUANTIZATION"),
|
||||
"dtype": "auto" if os.getenv("QUANTIZATION") is None else "half",
|
||||
"disable_log_stats": bool(int(os.getenv("DISABLE_LOG_STATS", 1))),
|
||||
"disable_log_requests": bool(int(os.getenv("DISABLE_LOG_REQUESTS", 1))),
|
||||
"gpu_memory_utilization": float(os.getenv("GPU_MEMORY_UTILIZATION", 0.98)),
|
||||
"tensor_parallel_size": self._get_num_gpu_shard(),
|
||||
}
|
||||
|
||||
def _initialize_llm(self):
|
||||
try:
|
||||
return AsyncLLMEngine.from_engine_args(AsyncEngineArgs(**self.config))
|
||||
except Exception as e:
|
||||
logging.error("Error initializing vLLM engine: %s", e)
|
||||
raise e
|
||||
|
||||
def _get_num_gpu_shard(self):
|
||||
final_num_gpu_shard = 1
|
||||
if bool(int(os.getenv("USE_TENSOR_PARALLEL", 0))):
|
||||
env_num_gpu_shard = int(os.getenv("TENSOR_PARALLEL_SIZE", 1))
|
||||
num_gpu_available = torch.cuda.device_count()
|
||||
final_num_gpu_shard = min(env_num_gpu_shard, num_gpu_available)
|
||||
logging.info("Using %s GPU shards", final_num_gpu_shard)
|
||||
return final_num_gpu_shard
|
||||
|
||||
def _get_n_current_jobs(self):
|
||||
total_sequences = len(self.llm.engine.scheduler.waiting) + len(self.llm.engine.scheduler.swapped) + len(self.llm.engine.scheduler.running)
|
||||
return total_sequences
|
||||
|
||||
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
|
||||
|
||||
+48
-37
@@ -1,56 +1,67 @@
|
||||
#!/usr/bin/env python
|
||||
from typing import Generator
|
||||
import runpod
|
||||
from utils import validate_and_convert_sampling_params, initialize_llm_engine, ServerlessConfig
|
||||
from vllm.utils import random_uuid
|
||||
|
||||
serverless_config = ServerlessConfig()
|
||||
llm, tokenizer = initialize_llm_engine()
|
||||
|
||||
def concurrency_modifier(current_concurrency) -> int:
|
||||
return serverless_config.max_concurrency
|
||||
from utils import validate_sampling_params, random_uuid
|
||||
from engine import VLLMEngine
|
||||
|
||||
vllm_engine = VLLMEngine()
|
||||
async def handler(job: dict) -> Generator[dict, None, None]:
|
||||
job_input = job["input"]
|
||||
prompt = job_input.get("prompt")
|
||||
llm_input = job_input.get("messages", job_input.get("prompt"))
|
||||
apply_chat_template = job_input.get("apply_chat_template", False)
|
||||
messages = job_input.get("messages")
|
||||
|
||||
if messages:
|
||||
prompt = tokenizer.apply_chat_template(messages)
|
||||
elif prompt and apply_chat_template:
|
||||
prompt = tokenizer.apply_chat_template(prompt)
|
||||
elif not prompt:
|
||||
raise ValueError("Must specify prompt or messages")
|
||||
|
||||
|
||||
if apply_chat_template 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", serverless_config.default_batch_size)
|
||||
batch_size = job_input.get("batch_size", vllm_engine.serverless_config.default_batch_size)
|
||||
sampling_params = job_input.get("sampling_params", {})
|
||||
|
||||
validated_params = validate_and_convert_sampling_params(sampling_params)
|
||||
validated_params = validate_sampling_params(sampling_params)
|
||||
request_id = random_uuid()
|
||||
results_generator = llm.generate(prompt, validated_params, request_id)
|
||||
results_generator = vllm_engine.llm.generate(
|
||||
llm_input, validated_params, request_id
|
||||
)
|
||||
|
||||
batch, last_output_text = [], ""
|
||||
batch = {"tokens": []}
|
||||
last_output_text = ""
|
||||
n_input_tokens, is_first_output = 0, True
|
||||
|
||||
async for request_output in results_generator:
|
||||
for output in request_output.outputs:
|
||||
usage = {"input": len(request_output.prompt_token_ids), "output": len(output.token_ids)}
|
||||
if is_first_output: # Count input tokens only once
|
||||
n_input_tokens = len(request_output.prompt_token_ids)
|
||||
is_first_output = False
|
||||
|
||||
for output in request_output.outputs:
|
||||
if stream:
|
||||
batch.append({"text": output.text[len(last_output_text):], "usage": usage})
|
||||
if len(batch) >= batch_size:
|
||||
|
||||
batch["tokens"].append(
|
||||
output.text[len(last_output_text):]
|
||||
)
|
||||
finished = request_output.finished
|
||||
if len(batch["tokens"]) >= batch_size or finished:
|
||||
batch["usage"] = {
|
||||
"input": n_input_tokens,
|
||||
"output": len(output.token_ids),
|
||||
}
|
||||
batch["finished"] = finished
|
||||
yield batch
|
||||
batch = []
|
||||
batch = {"tokens": []}
|
||||
|
||||
last_output_text = output.text
|
||||
|
||||
|
||||
if not stream:
|
||||
yield [{"text": last_output_text, "usage": usage}]
|
||||
yield {"tokens": [last_output_text],
|
||||
"usage": {
|
||||
"input": n_input_tokens,
|
||||
"output": len(output.token_ids),
|
||||
},
|
||||
"finished": True}
|
||||
|
||||
if batch:
|
||||
yield batch
|
||||
|
||||
runpod.serverless.start({
|
||||
"handler": handler,
|
||||
"concurrency_modifier": concurrency_modifier,
|
||||
"return_aggregate_stream": True
|
||||
})
|
||||
runpod.serverless.start(
|
||||
{
|
||||
"handler": handler,
|
||||
"concurrency_modifier": lambda x: vllm_engine.serverless_config.max_concurrency,
|
||||
"return_aggregate_stream": True,
|
||||
}
|
||||
)
|
||||
|
||||
+19
-60
@@ -1,18 +1,21 @@
|
||||
import os
|
||||
import logging
|
||||
from typing import Any, Dict, Optional, Union, Tuple
|
||||
from vllm import AsyncLLMEngine, AsyncEngineArgs, SamplingParams
|
||||
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 transformers import AutoTokenizer
|
||||
|
||||
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))
|
||||
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):
|
||||
@@ -22,47 +25,8 @@ class ServerlessConfig:
|
||||
def default_batch_size(self):
|
||||
return self._default_batch_size
|
||||
|
||||
class EngineConfig:
|
||||
def __init__(self):
|
||||
self.model_name = os.getenv('MODEL_NAME', 'default_model')
|
||||
self.tokenizer = os.getenv('TOKENIZER', self.model_name)
|
||||
self.model_base_path = os.getenv('MODEL_BASE_PATH', "/runpod-volume/")
|
||||
self.num_gpu_shard = int(os.getenv('NUM_GPU_SHARD', 1))
|
||||
self.use_full_metrics = os.getenv('USE_FULL_METRICS', 'True') == 'True'
|
||||
self.quantization = os.getenv('QUANTIZATION', None)
|
||||
self.dtype = "auto" if self.quantization is None else "half"
|
||||
self.disable_log_stats = os.getenv('DISABLE_LOG_STATS', 'True') == 'True'
|
||||
self.gpu_memory_utilization = float(os.getenv('GPU_MEMORY_UTILIZATION', 0.98))
|
||||
os.makedirs(self.model_base_path, exist_ok=True)
|
||||
|
||||
class Tokenizer:
|
||||
def __init__(self, tokenizer_name: str):
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
|
||||
|
||||
def apply_chat_template(self, input: Union[str, list[dict[str, str]]]) -> str:
|
||||
messages = input if isinstance(input, list) else [{"role": "user", "content": input}]
|
||||
return self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
||||
|
||||
def initialize_llm_engine() -> Tuple[AsyncLLMEngine, Tokenizer]:
|
||||
try:
|
||||
config = EngineConfig()
|
||||
engine_args = AsyncEngineArgs(
|
||||
model=config.model_name,
|
||||
download_dir=config.model_base_path,
|
||||
tokenizer=config.tokenizer,
|
||||
tensor_parallel_size=config.num_gpu_shard,
|
||||
dtype=config.dtype,
|
||||
disable_log_stats=config.disable_log_stats,
|
||||
quantization=config.quantization,
|
||||
gpu_memory_utilization=config.gpu_memory_utilization,
|
||||
)
|
||||
return AsyncLLMEngine.from_engine_args(engine_args), Tokenizer(config.tokenizer)
|
||||
except Exception as e:
|
||||
logging.error(f"Error initializing vLLM engine: {e}")
|
||||
raise
|
||||
|
||||
|
||||
def validate_and_convert_sampling_params(params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def validate_sampling_params(params: Dict[str, Any]) -> SamplingParams:
|
||||
validated_params = {}
|
||||
|
||||
for key, value in params.items():
|
||||
@@ -74,19 +38,14 @@ def validate_and_convert_sampling_params(params: Dict[str, Any]) -> Dict[str, An
|
||||
if expected_type is None:
|
||||
continue
|
||||
|
||||
if not isinstance(expected_type, tuple):
|
||||
expected_type = (expected_type,)
|
||||
|
||||
if any(isinstance(value, t) for t in expected_type):
|
||||
validated_params[key] = value
|
||||
if isinstance(expected_type, tuple):
|
||||
casted_value = next(
|
||||
(t(value) for t in expected_type if isinstance(value, t)), None
|
||||
)
|
||||
else:
|
||||
try:
|
||||
casted_value = next(
|
||||
t(value) for t in expected_type
|
||||
if isinstance(value, t)
|
||||
)
|
||||
validated_params[key] = casted_value
|
||||
except (TypeError, ValueError, StopIteration):
|
||||
continue
|
||||
casted_value = value if isinstance(value, expected_type) else None
|
||||
|
||||
return SamplingParams(**validated_params)
|
||||
if casted_value is not None:
|
||||
validated_params[key] = casted_value
|
||||
|
||||
return SamplingParams(**validated_params)
|
||||
|
||||
Reference in New Issue
Block a user