Small refactor

This commit is contained in:
alpayariyak
2024-03-06 17:08:57 +00:00
parent fee8d8eee4
commit c8ee100d80
6 changed files with 43 additions and 54 deletions
+3 -4
View File
@@ -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)
+1
View File
@@ -7,3 +7,4 @@ huggingface-hub
packaging
typing-extensions==4.7.1
pydantic
pydantic-settings
-26
View File
@@ -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
}
+5 -5
View File
@@ -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}
+31
View File
@@ -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())
+1 -17
View File
@@ -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