Chat Template
This commit is contained in:
+10
-2
@@ -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
@@ -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)
|
||||
Reference in New Issue
Block a user