From fe6e618692c3473d0267ff5600f8e4e56a7cf2ac Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Fri, 22 Dec 2023 04:08:55 +0000 Subject: [PATCH] Refactor code and Black formatting style --- src/constants.py | 42 +++++++++++------------ src/download_model.py | 8 +++-- src/engine.py | 59 ++++++++++++++++++++++++++++++++ src/handler.py | 59 ++++++++++++++++++-------------- src/utils.py | 79 +++++++++++-------------------------------- 5 files changed, 138 insertions(+), 109 deletions(-) create mode 100644 src/engine.py diff --git a/src/constants.py b/src/constants.py index c522e6a..e1866cc 100644 --- a/src/constants.py +++ b/src/constants.py @@ -2,24 +2,24 @@ DEFAULT_BATCH_SIZE = 10 DEFAULT_MAX_CONCURRENCY = 100 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, -} \ No newline at end of file + "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, +} diff --git a/src/download_model.py b/src/download_model.py index 3e9361f..999f903 100644 --- a/src/download_model.py +++ b/src/download_model.py @@ -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, - ) \ No newline at end of file + ) diff --git a/src/engine.py b/src/engine.py new file mode 100644 index 0000000..8b3d11e --- /dev/null +++ b/src/engine.py @@ -0,0 +1,59 @@ +import os +import logging +from typing import Union +import torch +from vllm import AsyncLLMEngine, AsyncEngineArgs +from transformers import AutoTokenizer + + +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) and 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): + self.config = self._initialize_config() + self.tokenizer = Tokenizer(self.config["model"]) + self.llm = self._initialize_llm() + + def _initialize_config(self): + return { + "model": os.getenv("MODEL_NAME", "default_model"), + "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))), + "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 diff --git a/src/handler.py b/src/handler.py index 0e2328d..5e95e5c 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,43 +1,49 @@ #!/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 +from utils import validate_sampling_params, ServerlessConfig, random_uuid +from engine import VLLMEngine + serverless_config = ServerlessConfig() -llm, tokenizer = initialize_llm_engine() +vllm_engine = VLLMEngine() + 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.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") - + llm_input, apply_chat_template = job_input.get( + "prompt", job_input["messages"] + ), job_input.get("apply_chat_template", False) + + 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) 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 = [], "" 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)} - + usage = { + "input": len(request_output.prompt_token_ids), + "output": len(output.token_ids), + } + if stream: - batch.append({"text": output.text[len(last_output_text):], "usage": usage}) + batch.append( + {"text": output.text[len(last_output_text) :], "usage": usage} + ) if len(batch) >= batch_size: yield batch batch = [] @@ -48,9 +54,12 @@ async def handler(job: dict) -> Generator[dict, None, None]: if batch: yield batch - -runpod.serverless.start({ - "handler": handler, - "concurrency_modifier": concurrency_modifier, - "return_aggregate_stream": True -}) + + +runpod.serverless.start( + { + "handler": handler, + "concurrency_modifier": concurrency_modifier, + "return_aggregate_stream": True, + } +) diff --git a/src/utils.py b/src/utils.py index 2979d6e..9445ec8 100644 --- a/src/utils.py +++ b/src/utils.py @@ -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) \ No newline at end of file + if casted_value is not None: + validated_params[key] = casted_value + + return SamplingParams(**validated_params)