This commit is contained in:
2026-03-29 16:49:09 +08:00
parent 5c6c59a362
commit 22afc63ac7
5 changed files with 31 additions and 101 deletions
+25 -28
View File
@@ -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,