x
This commit is contained in:
+25
-28
@@ -1,6 +1,4 @@
|
||||
from threading import Lock
|
||||
|
||||
from vllm import LLM, SamplingParams
|
||||
import httpx
|
||||
|
||||
from app.config import Settings
|
||||
from app.schemas import GenerateRequest, GenerateResponse
|
||||
@@ -9,36 +7,35 @@ 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,
|
||||
)
|
||||
self.client = httpx.Client(timeout=300.0)
|
||||
|
||||
def close(self) -> None:
|
||||
return None
|
||||
self.client.close()
|
||||
|
||||
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,
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.settings.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.settings.api_key}"
|
||||
payload = {
|
||||
"model": self.settings.served_model_name or self.settings.model_name,
|
||||
"messages": [{"role": "user", "content": req.prompt}],
|
||||
"max_tokens": req.max_tokens,
|
||||
"temperature": req.temperature,
|
||||
"top_p": req.top_p,
|
||||
}
|
||||
if req.stop:
|
||||
payload["stop"] = req.stop
|
||||
response = self.client.post(
|
||||
f"{self.settings.vllm_openai_internal_url}/chat/completions",
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
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)
|
||||
response.raise_for_status()
|
||||
body = response.json()
|
||||
completion = body["choices"][0]["message"]["content"]
|
||||
usage = body.get("usage", {})
|
||||
usage_prompt = int(usage.get("prompt_tokens", 0))
|
||||
usage_completion = int(usage.get("completion_tokens", 0))
|
||||
return GenerateResponse(
|
||||
text=completion,
|
||||
prompt=req.prompt,
|
||||
|
||||
Reference in New Issue
Block a user