From c8ee100d802e5211a109e5daabd2270b5bf91328 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Wed, 6 Mar 2024 17:08:57 +0000 Subject: [PATCH] Small refactor --- builder/download_model.py | 7 +++---- builder/requirements.txt | 3 ++- src/constants.py | 28 +--------------------------- src/engine.py | 10 +++++----- src/sampling_params.py | 31 +++++++++++++++++++++++++++++++ src/utils.py | 18 +----------------- 6 files changed, 43 insertions(+), 54 deletions(-) create mode 100644 src/sampling_params.py diff --git a/builder/download_model.py b/builder/download_model.py index dddcf56..4e1b783 100644 --- a/builder/download_model.py +++ b/builder/download_model.py @@ -45,7 +45,6 @@ if __name__ == "__main__": with open("/local_model_path.txt", "w") as f: f.write(model_folder) - if tokenizer != model: - tokenizer_folder = download_extras_or_tokenizer(tokenizer, download_dir, revisions["tokenizer"]) - with open("/local_tokenizer_path.txt", "w") as f: - f.write(tokenizer_folder) + tokenizer_folder = download_extras_or_tokenizer(tokenizer, download_dir, revisions["tokenizer"]) + with open("/local_tokenizer_path.txt", "w") as f: + f.write(tokenizer_folder) diff --git a/builder/requirements.txt b/builder/requirements.txt index 7e8cf28..39a3189 100644 --- a/builder/requirements.txt +++ b/builder/requirements.txt @@ -6,4 +6,5 @@ runpod==1.6.2 huggingface-hub packaging typing-extensions==4.7.1 -pydantic \ No newline at end of file +pydantic +pydantic-settings \ No newline at end of file diff --git a/src/constants.py b/src/constants.py index ce056b5..a75b5f1 100644 --- a/src/constants.py +++ b/src/constants.py @@ -1,30 +1,4 @@ -from typing import Union - DEFAULT_BATCH_SIZE = 50 DEFAULT_MAX_CONCURRENCY = 300 DEFAULT_BATCH_SIZE_GROWTH_FACTOR = 3 -DEFAULT_MIN_BATCH_SIZE = 1 - -SAMPLING_PARAM_TYPES = { - "n": int, - "best_of": int, - "presence_penalty": float, - "frequency_penalty": float, - "repetition_penalty": float, - "temperature": Union[float, int], - "top_p": float, - "top_k": int, - "min_p": float, - "use_beam_search": bool, - "length_penalty": float, - "early_stopping": Union[bool, str], - "stop": Union[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, - "include_stop_str_in_output": bool -} \ No newline at end of file +DEFAULT_MIN_BATCH_SIZE = 1 \ No newline at end of file diff --git a/src/engine.py b/src/engine.py index 9cfed73..a594a86 100644 --- a/src/engine.py +++ b/src/engine.py @@ -6,7 +6,7 @@ from dotenv import load_dotenv from torch.cuda import device_count from typing import AsyncGenerator -from vllm import AsyncLLMEngine, AsyncEngineArgs, SamplingParams +from vllm import AsyncLLMEngine, AsyncEngineArgs from vllm.entrypoints.openai.serving_chat import OpenAIServingChat from vllm.entrypoints.openai.serving_completion import OpenAIServingCompletion from vllm.entrypoints.openai.protocol import ChatCompletionRequest, CompletionRequest, ErrorResponse @@ -15,7 +15,7 @@ from utils import DummyRequest, JobInput, BatchSize, create_error_response from constants import DEFAULT_MAX_CONCURRENCY, DEFAULT_BATCH_SIZE, DEFAULT_BATCH_SIZE_GROWTH_FACTOR, DEFAULT_MIN_BATCH_SIZE from tokenizer import TokenizerWrapper from config import EngineConfig - +from sampling_params import validate_sampling_params class vLLMEngine: def __init__(self, engine = None): @@ -33,9 +33,10 @@ class vLLMEngine: async def generate(self, job_input: JobInput): try: + validated_sampling_params = validate_sampling_params(job_input.input_sampling_params) async for batch in self._generate_vllm( llm_input=job_input.llm_input, - validated_sampling_params=job_input.validated_sampling_params, + validated_sampling_params=validated_sampling_params, batch_size=job_input.max_batch_size, stream=job_input.stream, apply_chat_template=job_input.apply_chat_template, @@ -45,12 +46,11 @@ class vLLMEngine: ): yield batch except Exception as e: - yield create_error_response(str(e)).model_dump() + yield {"error": create_error_response(str(e)).model_dump()} async def _generate_vllm(self, llm_input, validated_sampling_params, batch_size, stream, apply_chat_template, request_id, batch_size_growth_factor, min_batch_size: str) -> AsyncGenerator[dict, None]: if apply_chat_template or isinstance(llm_input, list): llm_input = self.tokenizer.apply_chat_template(llm_input) - validated_sampling_params = SamplingParams(**validated_sampling_params) results_generator = self.llm.generate(llm_input, validated_sampling_params, request_id) n_responses, n_input_tokens, is_first_output = validated_sampling_params.n, 0, True last_output_texts, token_counters = ["" for _ in range(n_responses)], {"batch": 0, "total": 0} diff --git a/src/sampling_params.py b/src/sampling_params.py new file mode 100644 index 0000000..33fc26f --- /dev/null +++ b/src/sampling_params.py @@ -0,0 +1,31 @@ +from pydantic import BaseModel +from typing import Union, List, Optional +from vllm import SamplingParams + +class InputSamplingParams(BaseModel): + n: Optional[int] + best_of: Optional[int] + presence_penalty: Optional[float] + frequency_penalty: Optional[float] + repetition_penalty: Optional[float] + temperature: Optional[Union[float, int]] + top_p: Optional[float] + top_k: Optional[int] + min_p: Optional[float] + use_beam_search: Optional[bool] + length_penalty: Optional[float] + early_stopping: Optional[Union[bool, str]] + stop: Optional[Union[str, List[str]]] + stop_token_ids: Optional[List[int]] + ignore_eos: Optional[bool] + max_tokens: Optional[int] + logprobs: Optional[int] + prompt_logprobs: Optional[int] + skip_special_tokens: Optional[bool] + spaces_between_special_tokens: Optional[bool] + include_stop_str_in_output: Optional[bool] + +def validate_sampling_params(params: dict) -> SamplingParams: + cast_params = InputSamplingParams(**params) + return SamplingParams(**cast_params.model_dump()) + diff --git a/src/utils.py b/src/utils.py index 3b9fff7..363337e 100644 --- a/src/utils.py +++ b/src/utils.py @@ -1,11 +1,9 @@ import logging from http import HTTPStatus from typing import Any, Dict -from constants import SAMPLING_PARAM_TYPES from vllm.utils import random_uuid from vllm.entrypoints.openai.protocol import ErrorResponse - logging.basicConfig(level=logging.INFO) def count_physical_cores(): @@ -25,20 +23,6 @@ def count_physical_cores(): return len(cores) -def validate_sampling_params(params: Dict[str, Any]) -> Dict[str, Any]: - validated_params = {} - invalid_params = [] - for key, value in params.items(): - expected_type = SAMPLING_PARAM_TYPES.get(key) - if expected_type and isinstance(value, expected_type): - validated_params[key] = value - else: - invalid_params.append(key) - - if len(invalid_params) > 0: - logging.warning("Ignoring invalid sampling params: %s", invalid_params) - - return validated_params class JobInput: def __init__(self, job): @@ -47,7 +31,7 @@ class JobInput: self.max_batch_size = job.get("max_batch_size") self.apply_chat_template = job.get("apply_chat_template", False) self.use_openai_format = job.get("use_openai_format", False) - self.validated_sampling_params = validate_sampling_params(job.get("sampling_params", {})) + self.input_sampling_params = job.get("sampling_params", {}) self.request_id = random_uuid() batch_size_growth_factor = job.get("batch_size_growth_factor") self.batch_size_growth_factor = float(batch_size_growth_factor) if batch_size_growth_factor else None