49 lines
1.8 KiB
Python
49 lines
1.8 KiB
Python
import httpx
|
|
|
|
from app.config import Settings
|
|
from app.schemas import GenerateRequest, GenerateResponse
|
|
|
|
|
|
class InferenceEngine:
|
|
def __init__(self, settings: Settings) -> None:
|
|
self.settings = settings
|
|
self.client = httpx.Client(timeout=300.0)
|
|
|
|
def close(self) -> None:
|
|
self.client.close()
|
|
|
|
def generate(self, req: GenerateRequest) -> GenerateResponse:
|
|
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
|
|
if req.enable_thinking is not None:
|
|
payload["chat_template_kwargs"] = {"enable_thinking": req.enable_thinking}
|
|
response = self.client.post(
|
|
f"{self.settings.vllm_openai_internal_url}/chat/completions",
|
|
headers=headers,
|
|
json=payload,
|
|
)
|
|
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,
|
|
model=self.settings.served_model_name or self.settings.model_name,
|
|
usage_prompt_tokens=usage_prompt,
|
|
usage_completion_tokens=usage_completion,
|
|
usage_total_tokens=usage_prompt + usage_completion,
|
|
)
|