From 582e21f97fcd70fd545b3035d04ffd6036055ed2 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Wed, 20 Dec 2023 18:37:54 -0500 Subject: [PATCH] Chat Template --- src/handler.py | 12 ++++++++++-- src/utils.py | 21 ++++++++++++++++----- 2 files changed, 26 insertions(+), 7 deletions(-) diff --git a/src/handler.py b/src/handler.py index d58d0f1..7c59d7a 100644 --- a/src/handler.py +++ b/src/handler.py @@ -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", {}) diff --git a/src/utils.py b/src/utils.py index 949b63f..2979d6e 100644 --- a/src/utils.py +++ b/src/utils.py @@ -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) \ No newline at end of file