File size: 1,764 Bytes
4c45df7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 | """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
|