Small refactor
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -7,3 +7,4 @@ huggingface-hub
|
||||
packaging
|
||||
typing-extensions==4.7.1
|
||||
pydantic
|
||||
pydantic-settings
|
||||
@@ -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
@@ -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}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user