dewpoint / handler.py
valkayuh's picture
Dewpoint release
4c45df7
Raw History Blame Contribute Delete
1.76 kB
"""Hugging Face Inference Endpoints handler.
Deploy this repository as an Inference Endpoint and it serves Dewpoint with no
further code. Request body:
{"inputs": "so i said meet at three thirty tuesday what do you think",
"parameters": {"lang": "en"}}
`inputs` may also be a list of strings. Optional parameters:
lang ISO 639-1 code (default "en")
single true to run only the mmBERT-base member
labels true to also return per-word punctuation and case labels
Response: [{"text": "..."}] per input (plus "words", "punct", "case" if
labels=true).
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from dewpoint import Punctuator, split_words # noqa: E402
class EndpointHandler:
def __init__(self, path=""):
path = path or os.path.dirname(os.path.abspath(__file__))
self.full = Punctuator(path)
self._single = None
self.path = path
def _model(self, single):
if not single:
return self.full
if self._single is None:
self._single = Punctuator(self.path, members=["mmbert-base"])
return self._single
def __call__(self, data):
inputs = data.get("inputs", data)
params = data.get("parameters") or {}
lang = params.get("lang", "en")
p = self._model(bool(params.get("single", False)))
texts = [inputs] if isinstance(inputs, str) else list(inputs)
out = []
for t in texts:
item = {"text": p.restore(t, lang)}
if params.get("labels"):
r = p.predict(split_words(t, lang), lang)
item.update(words=r["words"], punct=r["punct"], case=r["case"])
out.append(item)
return out