dialingua-api / app.py
Lordkiki's picture
Upload app.py with huggingface_hub
b54d4ca verified
Raw
History Blame Contribute Delete
6.31 kB
#!/usr/bin/env python3
"""
Dialingua API — always-on translation service, small enough for a free tier.
WHY ONNX AND NOT PYTORCH
------------------------
Measured on the real model:
python + onnxruntime + numpy 47 MB
+ transformers (WITH torch) 390 MB <- torch alone is +343 MB
+ tokenizer 418 MB
+ encoder session 479 MB
+ decoder session 579 MB
+ decoder-with-past 672 MB
672 MB does not fit a 512 MB free tier. But torch is never used here — ONNX
Runtime does the inference and the tokenizer is pure sentencepiece. Leaving it
out of requirements.txt, and skipping the KV-cache session, brings this to
roughly 330 MB resident.
That is why requirements.txt pins `transformers` with NO torch. If torch ever
sneaks back in as a transitive dependency, this service will OOM on boot.
Generation is greedy and hand-rolled against the two ONNX sessions, because
optimum's generate() imports torch and would undo the whole point. Verses are
short, so the quadratic cost of re-running the decoder each step is cheap.
GET /health
POST /translate {"text": "..."}
POST /detect {"text": "..."}
"""
import os
import pathlib
import time
import numpy as np
import onnxruntime as ort
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
from transformers import AutoTokenizer
MODEL_ID = os.environ.get("BKV_MODEL", "Lordkiki/dialingua-bkv2eng-web")
MAX_NEW = int(os.environ.get("BKV_MAX_TOKENS", "160"))
ORIGINS = [o.strip() for o in os.environ.get("ALLOWED_ORIGINS", "*").split(",")]
_state = {}
app = FastAPI(title="Dialingua API", version="1.0.0")
app.add_middleware(CORSMiddleware, allow_origins=ORIGINS,
allow_credentials=False, allow_methods=["*"],
allow_headers=["*"])
class TextIn(BaseModel):
text: str = Field(min_length=1, max_length=2000)
def _session(path: str) -> ort.InferenceSession:
opts = ort.SessionOptions()
# One thread: free tiers give a fraction of a core, and extra threads cost
# memory without buying speed.
opts.intra_op_num_threads = 1
opts.inter_op_num_threads = 1
opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
return ort.InferenceSession(path, opts, providers=["CPUExecutionProvider"])
@app.on_event("startup")
def load() -> None:
from huggingface_hub import snapshot_download
t0 = time.perf_counter()
# Only the two graphs greedy decoding needs. Pulling the KV-cache decoder
# too would add ~93 MB of resident memory for no benefit here.
local = snapshot_download(
MODEL_ID,
allow_patterns=["*.json", "*.spm", "*.model",
"onnx/encoder_model_quantized.onnx",
"onnx/decoder_model_quantized.onnx"],
)
root = pathlib.Path(local)
_state["tok"] = AutoTokenizer.from_pretrained(local)
_state["enc"] = _session(str(root / "onnx" / "encoder_model_quantized.onnx"))
_state["dec"] = _session(str(root / "onnx" / "decoder_model_quantized.onnx"))
_state["dec_inputs"] = {i.name for i in _state["dec"].get_inputs()}
print(f"ready in {time.perf_counter() - t0:.0f}s ({MODEL_ID})")
@app.get("/health")
def health():
return {"status": "ok", "model": MODEL_ID, "loaded": "enc" in _state}
@app.post("/translate")
def translate(body: TextIn):
if "enc" not in _state:
raise HTTPException(503, "model still loading")
tok, enc, dec = _state["tok"], _state["enc"], _state["dec"]
ids = tok(body.text, return_tensors="np", truncation=True, max_length=256)
input_ids = ids["input_ids"].astype(np.int64)
attention = ids["attention_mask"].astype(np.int64)
hidden = enc.run(None, {"input_ids": input_ids,
"attention_mask": attention})[0]
# Marian starts decoding from pad_token_id.
start = tok.pad_token_id if tok.pad_token_id is not None else 0
eos = tok.eos_token_id
out_ids = [start]
for _ in range(MAX_NEW):
feed = {"encoder_attention_mask": attention,
"encoder_hidden_states": hidden,
"input_ids": np.array([out_ids], dtype=np.int64)}
feed = {k: v for k, v in feed.items() if k in _state["dec_inputs"]}
logits = dec.run(None, feed)[0]
nxt = int(np.argmax(logits[0, -1]))
if nxt == eos:
break
out_ids.append(nxt)
text = tok.decode(out_ids[1:], skip_special_tokens=True).strip()
return {
"translation": text,
"direction": "bkv2eng",
"caveat": "Trained on ~1,100 scripture verse pairs. Formal register is "
"reasonable; everyday speech is not. Have a speaker check "
"anything that matters.",
}
@app.post("/detect")
def detect(body: TextIn):
"""Bekwarra detection, no model required.
Keys on the phonemic apostrophe (k'uchu, ng'amin — a letter here, not
punctuation), the kp/gb clusters common to Niger-Congo, and the high rate
of vowel-initial words. English contractions are subtracted so don't/it's
do not read as Bekwarra.
"""
import re
text = body.text
words = re.findall(r"[^\W\d_]+", text.lower(), flags=re.UNICODE)
if not words:
return {"code": None, "name": "—", "confidence": 0.0, "reason": "no words"}
apo = len(re.findall(
r"\b(?:ng|kp|gb|ch|sh|[bcdfghjklmnprstvwyz])'\s?[aeiou]", text, re.I))
contractions = len(re.findall(r"\b\w+'(?:s|t|re|ve|ll|d|m)\b", text, re.I))
apo = max(0, apo - contractions)
dig = len(re.findall(r"kp|gb", text, re.I))
vowel = sum(1 for w in words if w[:1] in "aeiou")
score = (min(apo / max(len(words) * .18, 1), 1) * .5
+ min(dig / max(len(words) * .10, 1), 1) * .2
+ min(vowel / len(words) / .4, 1) * .3)
if score > .42:
return {"code": "bkv", "name": "Bekwarra",
"confidence": round(min(.5 + score * .5, .99), 2),
"reason": f"{apo} phonemic apostrophes, {dig} kp/gb clusters"}
return {"code": None, "name": "Not Bekwarra",
"confidence": round(1 - score, 2),
"reason": "no Bekwarra orthographic signal"}