FastAPI service wrapping the Surya-OCR-2 model (datalab-to) served through vLLM: legacy /v1/api/ai/* endpoints, an OpenAI-compatible /v1/chat/completions endpoint, a coalescing request batcher, a local OCR CLI, Docker packaging, multilingual example outputs, and quantization/concurrency benchmarks. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
258 lines
8.2 KiB
Python
258 lines
8.2 KiB
Python
"""Shared OpenAI-compatible chat completions client. Used by vllm + llama.cpp.
|
|
|
|
Both servers expose `/v1/chat/completions` with the same request/response shape,
|
|
so this module is the single point of HTTP contact for both backends.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import io
|
|
import math
|
|
import os
|
|
import time
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from typing import List, Optional
|
|
|
|
from PIL import Image
|
|
|
|
from surya.inference.prompts import PROMPT_MAPPING
|
|
from surya.inference.schema import (
|
|
BatchInputItem,
|
|
BatchOutputItem,
|
|
GenerationResult,
|
|
)
|
|
from surya.inference.util import detect_repeat_token, scale_to_fit
|
|
from surya.logging import get_logger
|
|
from surya.timing import get_current_timing, set_current_timing, timing_span
|
|
|
|
logger = get_logger()
|
|
|
|
|
|
def resolve_max_workers(batch_len: int, max_inflight: int) -> int:
|
|
"""Concurrent HTTP workers for a batch: as many as the batch needs, capped
|
|
by max_inflight so we keep vLLM's sequence slots full without over-queueing."""
|
|
return max(1, min(batch_len, max_inflight))
|
|
|
|
|
|
def encode_image_b64(image: Image.Image) -> tuple[str, str]:
|
|
image_format = os.getenv("SUYA_VLLM_IMAGE_FORMAT", "JPEG").upper()
|
|
if image_format not in {"JPEG", "PNG"}:
|
|
raise ValueError("SUYA_VLLM_IMAGE_FORMAT must be JPEG or PNG")
|
|
|
|
buf = io.BytesIO()
|
|
if image_format == "JPEG":
|
|
quality = int(os.getenv("SUYA_VLLM_JPEG_QUALITY", "92"))
|
|
image.save(buf, format="JPEG", quality=quality, subsampling=0)
|
|
mime_type = "image/jpeg"
|
|
else:
|
|
image.save(buf, format="PNG")
|
|
mime_type = "image/png"
|
|
view = buf.getbuffer()
|
|
try:
|
|
return base64.b64encode(view).decode("ascii"), mime_type
|
|
finally:
|
|
view.release()
|
|
|
|
|
|
def _build_messages(image: Image.Image, prompt: str):
|
|
with timing_span("openai_encode_image_b64", image_size=image.size):
|
|
image_b64, mime_type = encode_image_b64(image)
|
|
return [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {"url": f"data:{mime_type};base64,{image_b64}"},
|
|
},
|
|
{"type": "text", "text": prompt},
|
|
],
|
|
}
|
|
]
|
|
|
|
|
|
def _mean_token_prob(logprobs_content) -> Optional[float]:
|
|
if not logprobs_content:
|
|
return None
|
|
probs = []
|
|
for tok in logprobs_content:
|
|
lp = (
|
|
tok.get("logprob")
|
|
if isinstance(tok, dict)
|
|
else getattr(tok, "logprob", None)
|
|
)
|
|
if lp is None:
|
|
continue
|
|
probs.append(math.exp(lp))
|
|
if not probs:
|
|
return None
|
|
return sum(probs) / len(probs)
|
|
|
|
|
|
def _generate_one(
|
|
item: BatchInputItem,
|
|
client,
|
|
model_name: str,
|
|
max_tokens_default: int,
|
|
temperature: float,
|
|
top_p: float,
|
|
timeout: float,
|
|
request_logprobs_default: bool,
|
|
) -> GenerationResult:
|
|
prompt = item.prompt or PROMPT_MAPPING[item.prompt_type]
|
|
with timing_span("openai_scale_image", prompt_type=item.prompt_type, image_size=item.image.size):
|
|
image = scale_to_fit(item.image)
|
|
with timing_span("openai_build_messages", prompt_type=item.prompt_type):
|
|
messages = _build_messages(image, prompt)
|
|
|
|
max_tokens = item.max_tokens or max_tokens_default
|
|
request_logprobs = item.request_logprobs or request_logprobs_default
|
|
|
|
kwargs = dict(
|
|
model=model_name,
|
|
messages=messages,
|
|
max_tokens=max_tokens,
|
|
temperature=temperature,
|
|
top_p=top_p,
|
|
timeout=timeout,
|
|
)
|
|
if request_logprobs:
|
|
kwargs["logprobs"] = True
|
|
|
|
# Structured output: prefer OpenAI-standard response_format (works on both
|
|
# vllm and llama.cpp). Fall back to vllm's extra_body for guided_regex.
|
|
if item.guided_json is not None:
|
|
kwargs["response_format"] = {
|
|
"type": "json_schema",
|
|
"json_schema": {
|
|
"name": "structured_output",
|
|
"schema": item.guided_json,
|
|
"strict": True,
|
|
},
|
|
}
|
|
if item.guided_regex is not None:
|
|
kwargs.setdefault("extra_body", {})["guided_regex"] = item.guided_regex
|
|
|
|
try:
|
|
with timing_span(
|
|
"openai_chat_completion",
|
|
prompt_type=item.prompt_type,
|
|
max_tokens=max_tokens,
|
|
):
|
|
completion = client.chat.completions.create(**kwargs)
|
|
raw = completion.choices[0].message.content or ""
|
|
token_count = completion.usage.completion_tokens if completion.usage else 0
|
|
with timing_span(
|
|
"openai_parse_completion",
|
|
prompt_type=item.prompt_type,
|
|
token_count=token_count,
|
|
):
|
|
mean_p = None
|
|
logprobs_content = None
|
|
if request_logprobs:
|
|
choice = completion.choices[0]
|
|
lp = getattr(choice, "logprobs", None)
|
|
if lp is not None:
|
|
content = getattr(lp, "content", None)
|
|
if content is not None:
|
|
logprobs_content = [
|
|
c.model_dump() if hasattr(c, "model_dump") else c
|
|
for c in content
|
|
]
|
|
mean_p = _mean_token_prob(content)
|
|
return GenerationResult(
|
|
raw=raw,
|
|
token_count=token_count,
|
|
error=False,
|
|
mean_token_prob=mean_p,
|
|
logprobs=logprobs_content,
|
|
)
|
|
except Exception as e:
|
|
logger.warning(f"Inference error: {e}")
|
|
return GenerationResult(raw="", token_count=0, error=True)
|
|
|
|
|
|
def _should_retry(
|
|
result: GenerationResult,
|
|
retries: int,
|
|
max_retries: int,
|
|
) -> bool:
|
|
if retries >= max_retries:
|
|
return False
|
|
if result.error:
|
|
return True
|
|
has_repeat = detect_repeat_token(result.raw) or (
|
|
len(result.raw) > 50 and detect_repeat_token(result.raw, cut_from_end=50)
|
|
)
|
|
return has_repeat
|
|
|
|
|
|
def chat_completions_batch(
|
|
batch: List[BatchInputItem],
|
|
client,
|
|
model_name: str,
|
|
max_tokens_default: int = 2048,
|
|
temperature: float = 0.0,
|
|
top_p: float = 0.1,
|
|
timeout: float = 600.0,
|
|
max_workers: Optional[int] = None,
|
|
max_retries: int = 3,
|
|
request_logprobs_default: bool = True,
|
|
) -> List[BatchOutputItem]:
|
|
"""Run a batch of items through the chat completions endpoint with concurrent workers."""
|
|
if not batch:
|
|
return []
|
|
if max_workers is None:
|
|
max_workers = min(64, len(batch))
|
|
collector = get_current_timing()
|
|
|
|
def _process(item: BatchInputItem) -> BatchOutputItem:
|
|
if collector is not None:
|
|
set_current_timing(collector)
|
|
result = _generate_one(
|
|
item,
|
|
client=client,
|
|
model_name=model_name,
|
|
max_tokens_default=max_tokens_default,
|
|
temperature=temperature,
|
|
top_p=top_p,
|
|
timeout=timeout,
|
|
request_logprobs_default=request_logprobs_default,
|
|
)
|
|
retries = 0
|
|
while _should_retry(result, retries, max_retries):
|
|
backoff = 1.5 * (retries + 1) if result.error else 0
|
|
if backoff:
|
|
time.sleep(backoff)
|
|
retry_temp = min(temperature + 0.2 * (retries + 1), 0.8)
|
|
retry_top_p = 0.95 if not result.error else top_p
|
|
result = _generate_one(
|
|
item,
|
|
client=client,
|
|
model_name=model_name,
|
|
max_tokens_default=max_tokens_default,
|
|
temperature=retry_temp,
|
|
top_p=retry_top_p,
|
|
timeout=timeout,
|
|
request_logprobs_default=request_logprobs_default,
|
|
)
|
|
retries += 1
|
|
return BatchOutputItem(
|
|
raw=result.raw,
|
|
token_count=result.token_count,
|
|
error=result.error,
|
|
mean_token_prob=result.mean_token_prob,
|
|
logprobs=result.logprobs,
|
|
metadata=item.metadata,
|
|
)
|
|
|
|
with timing_span(
|
|
"openai_batch_threadpool",
|
|
item_count=len(batch),
|
|
max_workers=max_workers,
|
|
max_retries=max_retries,
|
|
):
|
|
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
|
return list(executor.map(_process, batch))
|