47 lines
1.7 KiB
Python
47 lines
1.7 KiB
Python
from threading import Lock
|
|
|
|
from vllm import LLM, SamplingParams
|
|
|
|
from app.config import Settings
|
|
from app.schemas import GenerateRequest, GenerateResponse
|
|
|
|
|
|
class InferenceEngine:
|
|
def __init__(self, settings: Settings) -> None:
|
|
self.settings = settings
|
|
self.lock = Lock()
|
|
self.model = LLM(
|
|
model=settings.model_name,
|
|
tensor_parallel_size=settings.tensor_parallel_size,
|
|
max_model_len=settings.max_model_len,
|
|
gpu_memory_utilization=settings.gpu_memory_utilization,
|
|
max_num_seqs=settings.max_num_seqs,
|
|
dtype=settings.dtype,
|
|
enforce_eager=settings.enforce_eager,
|
|
trust_remote_code=settings.trust_remote_code,
|
|
revision=settings.revision,
|
|
)
|
|
|
|
def generate(self, req: GenerateRequest) -> GenerateResponse:
|
|
sampling_params = SamplingParams(
|
|
temperature=req.temperature,
|
|
top_p=req.top_p,
|
|
max_tokens=req.max_tokens,
|
|
repetition_penalty=req.repetition_penalty,
|
|
stop=req.stop,
|
|
)
|
|
with self.lock:
|
|
outputs = self.model.generate([req.prompt], sampling_params, use_tqdm=False)
|
|
output = outputs[0]
|
|
completion = output.outputs[0].text
|
|
usage_prompt = len(output.prompt_token_ids)
|
|
usage_completion = len(output.outputs[0].token_ids)
|
|
return GenerateResponse(
|
|
text=completion,
|
|
prompt=req.prompt,
|
|
model=self.settings.model_name,
|
|
usage_prompt_tokens=usage_prompt,
|
|
usage_completion_tokens=usage_completion,
|
|
usage_total_tokens=usage_prompt + usage_completion,
|
|
)
|