Refactor code and Black formatting style
This commit is contained in:
+21
-21
@@ -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,
|
||||
}
|
||||
"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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
+34
-25
@@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
+19
-60
@@ -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)
|
||||
if casted_value is not None:
|
||||
validated_params[key] = casted_value
|
||||
|
||||
return SamplingParams(**validated_params)
|
||||
|
||||
Reference in New Issue
Block a user