Chat Template

This commit is contained in:
alpayariyak
2023-12-20 18:37:54 -05:00
parent dad02e13d8
commit 582e21f97f
2 changed files with 26 additions and 7 deletions
+10 -2
View File
@@ -5,14 +5,22 @@ from utils import validate_and_convert_sampling_params, initialize_llm_engine, S
from vllm.utils import random_uuid
serverless_config = ServerlessConfig()
llm = initialize_llm_engine()
llm, tokenizer = initialize_llm_engine()
def concurrency_modifier(current_concurrency) -> int:
return max(0, serverless_config.max_concurrency - current_concurrency)
async def handler(job: dict) -> Generator[dict, None, None]:
job_input = job["input"]
prompt = job_input["prompt"]
prompt = 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 apply_chat_template:
prompt = tokenizer.apply_chat_template(prompt)
streaming = job_input.get("streaming", False)
batch_size = job_input.get("batch_size", serverless_config.default_batch_size)
sampling_params = job_input.get("sampling_params", {})
+16 -5
View File
@@ -1,11 +1,14 @@
import os
from typing import Any, Dict, Optional, Union
import logging
from typing import Any, Dict, Optional, Union, Tuple
from vllm import AsyncLLMEngine, AsyncEngineArgs, SamplingParams
from constants import sampling_param_types, DEFAULT_BATCH_SIZE, DEFAULT_MAX_CONCURRENCY
import logging
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))
@@ -32,8 +35,15 @@ class EngineConfig:
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() -> AsyncLLMEngine:
def initialize_llm_engine() -> Tuple[AsyncLLMEngine, Tokenizer]:
try:
config = EngineConfig()
engine_args = AsyncEngineArgs(
@@ -46,10 +56,11 @@ def initialize_llm_engine() -> AsyncLLMEngine:
quantization=config.quantization,
gpu_memory_utilization=config.gpu_memory_utilization,
)
return AsyncLLMEngine.from_engine_args(engine_args)
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]:
validated_params = {}
@@ -78,4 +89,4 @@ def validate_and_convert_sampling_params(params: Dict[str, Any]) -> Dict[str, An
except (TypeError, ValueError, StopIteration):
continue
return SamplingParams(**validated_params)
return SamplingParams(**validated_params)