"""Dewpoint: multilingual punctuation restoration and truecasing. Self-contained inference for the released checkpoints, in one file, with two interchangeable backends: from dewpoint import Punctuator p = Punctuator.from_pretrained("valkayuh/dewpoint") p.restore("so i said meet at three thirty tuesday what do you think", lang="en") # -> 'So I said meet at three thirty Tuesday. What do you think?' backend="torch" needs torch + transformers + safetensors, and uses a GPU if present. backend="onnx" needs only onnxruntime + tokenizers + numpy (no torch at all), and downloads the ONNX graphs instead of the safetensors. backend="auto" (the default) uses torch when it is installed, otherwise ONNX. The model is an ensemble of two dual-head token taggers (mmBERT-base and XLM-RoBERTa-large). Each member reads the same word list and returns one posterior per word; the posteriors are averaged in probability space, a per-class decision bias fitted on IWSLT dev2010 is applied, and the argmax is taken. Pass ``members=["mmbert-base"]`` for the single 307M-parameter model, which uses its own calibration. It is a tagger, not a generator: it never adds, drops, reorders or rewrites a word. It only decides, per word, which mark follows it (none / , / . / ? / !) and how it is cased (lower / Capitalised / UPPER). """ from __future__ import annotations import json import os import re import unicodedata from typing import Iterable, Sequence import numpy as np try: # torch is optional: the ONNX backend needs none import torch import torch.nn as nn except ImportError: # pragma: no cover torch = None nn = None __all__ = ["Punctuator", "StreamingPunctuator"] PUNCT_LABELS = ["O", "COMMA", "PERIOD", "QUESTION", "EXCLAM"] CASE_LABELS = ["LOWER", "CAP", "UPPER"] _SENT_END = {"PERIOD", "QUESTION", "EXCLAM"} # ------------------------------------------------------------ language rules # Scripts with no case distinction: the case head is masked to LOWER. CASELESS_LANGS = { "zh", "zh-cn", "zh-tw", "ja", "ko", "ar", "he", "fa", "ur", "th", "hi", "bn", "ta", "te", "ml", "kn", "mr", "gu", "pa", "ne", "si", "my", "km", "lo", "am", "ti", "dv", "bo", "dz", "ps", "sd", "ug", "yi", "yue", "wuu", "as", "or", "arq", "arz", "ary", "acm", "apc", "ckb", "ka", "he-il", "sa", "ks", "bho", } # Written without spaces between words: one grapheme cluster is one token. NO_SPACE_LANGS = {"zh", "zh-cn", "zh-tw", "ja", "th", "my", "km", "lo", "bo", "dz", "yue", "wuu"} _SURFACE = { "ascii": {"O": "", "COMMA": ",", "PERIOD": ".", "QUESTION": "?", "EXCLAM": "!"}, "cjk": {"O": "", "COMMA": ",", "PERIOD": "。", "QUESTION": "?", "EXCLAM": "!"}, "ja": {"O": "", "COMMA": "、", "PERIOD": "。", "QUESTION": "?", "EXCLAM": "!"}, "arab": {"O": "", "COMMA": "،", "PERIOD": ".", "QUESTION": "؟", "EXCLAM": "!"}, "deva": {"O": "", "COMMA": ",", "PERIOD": "।", "QUESTION": "?", "EXCLAM": "!"}, # Greek writes its question mark with the ASCII semicolon. "el": {"O": "", "COMMA": ",", "PERIOD": ".", "QUESTION": ";", "EXCLAM": "!"}, "hy": {"O": "", "COMMA": ",", "PERIOD": "։", "QUESTION": "՞", "EXCLAM": "՜"}, } _SURFACE_BY_LANG = {} for _l in ("zh", "zh-cn", "zh-tw", "yue", "wuu"): _SURFACE_BY_LANG[_l] = _SURFACE["cjk"] for _l in ("ar", "fa", "ur", "ps", "arq", "arz", "ary", "ckb", "sd", "ug"): _SURFACE_BY_LANG[_l] = _SURFACE["arab"] for _l in ("hi", "bn", "mr", "ne", "pa", "as", "or", "sa", "bho"): _SURFACE_BY_LANG[_l] = _SURFACE["deva"] _SURFACE_BY_LANG.update(el=_SURFACE["el"], hy=_SURFACE["hy"], ja=_SURFACE["ja"]) # Characters stripped from the edges of an input token before tagging, so text # that already carries some punctuation is handled the same as bare ASR output. _MARKS = set( ".。.۔।॥။։។…⋯።,,،٫၊、᠂᠈፣??؟;՞⸮፧!!՜::;؛·-–—―−፡" "\"'‘’“”„‚«»()[]{}〈〉《》「」『』()‹›‟″′¡¿⸘‛*_`•⁃・"'´‐‑" ) _LATIN_RUN = re.compile(r"[A-Za-z0-9À-ɏ]+(?:[.'’-][A-Za-z0-9]+)*") def base_lang(lang: str) -> str: return lang.split("-")[0].lower() def is_caseless(lang: str) -> bool: return lang.lower() in CASELESS_LANGS or base_lang(lang) in CASELESS_LANGS def is_no_space(lang: str) -> bool: return lang.lower() in NO_SPACE_LANGS or base_lang(lang) in NO_SPACE_LANGS def surface_for(lang: str) -> dict: return _SURFACE_BY_LANG.get(base_lang(lang), _SURFACE["ascii"]) def apply_case(word: str, case: str) -> str: if case == "UPPER": return word.upper() if case == "CAP": return word[:1].upper() + word[1:] return word def split_words(text: str, lang: str) -> list[str]: """Raw text -> the lowercased, mark-free word list the model was trained on.""" text = unicodedata.normalize("NFC", text) if is_no_space(lang): return _split_charwise(text) out = [] for raw in text.split(): chars = [i for i, ch in enumerate(raw) if ch not in _MARKS] if not chars: continue core = raw[chars[0]:chars[-1] + 1] if any(ch.isalnum() for ch in core): out.append(core.lower()) return out def _split_charwise(text: str) -> list[str]: """CJK / Thai / Khmer: grapheme clusters, with Latin runs kept whole.""" words, i, n = [], 0, len(text) while i < n: ch = text[i] if ch.isspace() or ch in _MARKS: i += 1 continue m = _LATIN_RUN.match(text, i) if m and m.group(): words.append(m.group().lower()) i = m.end() continue if ch.isalnum() or unicodedata.category(ch) in ("Mn", "Mc", "Me"): j = i + 1 while j < n and unicodedata.category(text[j]) in ("Mn", "Mc", "Me"): j += 1 words.append(text[i:j].lower()) i = j continue i += 1 return words # ------------------------------------------------------------------- model if torch is not None: class PunctCaseModel(nn.Module): """Encoder -> residual shared trunk -> punctuation head and case head.""" def __init__(self, config, n_punct=5, n_case=3): super().__init__() from transformers import AutoModel self.encoder = AutoModel.from_config(config) h = config.hidden_size self.drop = nn.Dropout(0.1) self.trunk = nn.Sequential(nn.Linear(h, h), nn.GELU(), nn.LayerNorm(h, eps=1e-5)) self.head_punct = nn.Linear(h, n_punct) self.head_case = nn.Linear(h, n_case) def forward(self, input_ids, attention_mask=None): x = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state x = self.drop(x) x = x + self.trunk(x) return self.head_punct(x), self.head_case(x) __all__.append("PunctCaseModel") def plan_windows(sub_counts, max_sub=508, overlap_sub=256): """Overlapping windows under a subword budget; only each window's centre commits. Returns (win_start, win_end, commit_start, commit_end) in words.""" n = len(sub_counts) if n == 0: return [] overlap_sub = min(overlap_sub, max_sub // 2) prefix = [0] * (n + 1) for i, c in enumerate(sub_counts): prefix[i + 1] = prefix[i] + max(1, c) def end_for(start): lo, hi = start + 1, n budget = prefix[start] + max_sub while lo < hi: mid = (lo + hi + 1) // 2 if prefix[mid] <= budget: lo = mid else: hi = mid - 1 return max(lo, start + 1) wins, start = [], 0 while True: end = end_for(start) wins.append([start, end]) if end >= n: break target = prefix[end] - overlap_sub nxt = end while nxt > start + 1 and prefix[nxt] > target: nxt -= 1 start = max(nxt, start + 1) out = [] for i, (s, e) in enumerate(wins): cs = s if i == 0 else (wins[i - 1][1] + s) // 2 ce = e if i == len(wins) - 1 else (e + wins[i + 1][0] + 1) // 2 cs = max(cs, s) out.append((s, e, cs, min(max(ce, cs), e))) return out class _MemberBase: """Windowing, first-subword gathering and caseless masking, shared by both backends so that they cannot drift apart. Subclasses supply `_counts` (subwords per word), `_encode` (a padded batch plus word ids) and `_forward` (softmaxed posteriors as numpy).""" def _load_meta(self, path): with open(os.path.join(path, "vpunct_config.json"), encoding="utf-8") as f: self.meta = json.load(f) self.max_len = int(self.meta.get("max_len", 512)) def posteriors(self, words, lang): n = len(words) pp = np.zeros((n, len(PUNCT_LABELS)), dtype=np.float32) pp[:, 0] = 1.0 cp = np.zeros((n, len(CASE_LABELS)), dtype=np.float32) cp[:, 0] = 1.0 if n == 0: return pp, cp wins = plan_windows(self._counts(list(words)), max_sub=self.max_len - 4, overlap_sub=256) for i in range(0, len(wins), self.batch_size): chunk = wins[i:i + self.batch_size] ids, mask, wids = self._encode([list(words[s:e]) for s, e, _, _ in chunk]) p, c = self._forward(ids, mask) for b, (s, e, cs, ce) in enumerate(chunk): loc_p = np.zeros((e - s, len(PUNCT_LABELS)), dtype=np.float32) loc_p[:, 0] = 1.0 loc_c = np.zeros((e - s, len(CASE_LABELS)), dtype=np.float32) loc_c[:, 0] = 1.0 seen = set() for pos, w in enumerate(wids[b]): if w is None or w in seen or pos >= p.shape[1]: continue seen.add(w) # label lives on a word's first subword loc_p[w] = p[b, pos] loc_c[w] = c[b, pos] pp[cs:ce] = loc_p[cs - s:ce - s] cp[cs:ce] = loc_c[cs - s:ce - s] if is_caseless(lang): cp[:] = 0.0 cp[:, 0] = 1.0 return pp, cp def _softmax(x): x = x - x.max(-1, keepdims=True) e = np.exp(x) return e / e.sum(-1, keepdims=True) class _TorchMember(_MemberBase): """One tagger run by torch, from its safetensors.""" def __init__(self, path, device, dtype, batch_size=64): from safetensors.torch import load_file from transformers import AutoConfig, AutoTokenizer self._load_meta(path) cfg = AutoConfig.from_pretrained(path) model = PunctCaseModel(cfg, self.meta["n_punct"], self.meta["n_case"]) model.load_state_dict(load_file(os.path.join(path, "model.safetensors"))) self.model = model.to(device=device, dtype=dtype).eval() try: self.model.encoder.config._attn_implementation = "sdpa" except Exception: pass self.tok = AutoTokenizer.from_pretrained(path) self.device = device self.batch_size = batch_size def _counts(self, words): return [len(e) for e in self.tok(words, add_special_tokens=False)["input_ids"]] def _encode(self, batch): enc = self.tok(batch, is_split_into_words=True, truncation=True, max_length=self.max_len, padding=True, return_tensors="pt") return enc["input_ids"], enc["attention_mask"], \ [enc.word_ids(b) for b in range(len(batch))] def _forward(self, ids, mask): with torch.no_grad(): lp, lc = self.model(ids.to(self.device), mask.to(self.device)) return (lp.float().softmax(-1).cpu().numpy(), lc.float().softmax(-1).cpu().numpy()) class _OnnxMember(_MemberBase): """One tagger run by onnxruntime, tokenised by the `tokenizers` library. No torch and no transformers anywhere on this path.""" def __init__(self, path, onnx_dir, providers=None, batch_size=8, threads=None): import onnxruntime as ort from tokenizers import Tokenizer self._load_meta(path) self.tok = Tokenizer.from_file(os.path.join(path, "tokenizer.json")) self.tok.no_padding() # truncate the way transformers does: before the special tokens are # added, so an over-long window still ends in its closing token self.tok.enable_truncation(max_length=self.max_len) with open(os.path.join(path, "tokenizer_config.json"), encoding="utf-8") as f: tcfg = json.load(f) self.pad_id = self.tok.token_to_id(tcfg.get("pad_token", "")) or 0 so = ort.SessionOptions() if threads: so.intra_op_num_threads = threads if providers is None: avail = ort.get_available_providers() providers = [p for p in ("CUDAExecutionProvider", "CoreMLExecutionProvider", "DmlExecutionProvider") if p in avail] providers.append("CPUExecutionProvider") self.sess = ort.InferenceSession(os.path.join(onnx_dir, "model.onnx"), so, providers=providers) self.batch_size = batch_size def _counts(self, words): return [len(e.ids) for e in self.tok.encode_batch(words, add_special_tokens=False)] def _encode(self, batch): encs = self.tok.encode_batch(batch, is_pretokenized=True) L = min(self.max_len, max(len(e.ids) for e in encs)) ids = np.full((len(encs), L), self.pad_id, dtype=np.int64) mask = np.zeros((len(encs), L), dtype=np.int64) wids = [] for b, e in enumerate(encs): n = min(L, len(e.ids)) ids[b, :n] = e.ids[:n] mask[b, :n] = 1 wids.append(list(e.word_ids[:n])) return ids, mask, wids def _forward(self, ids, mask): lp, lc = self.sess.run(None, {"input_ids": ids, "attention_mask": mask}) return _softmax(lp.astype(np.float32)), _softmax(lc.astype(np.float32)) # -------------------------------------------------------------- public api class Punctuator: """Restore punctuation and case in unpunctuated, lowercased text.""" MEMBERS = ("mmbert-base", "xlm-roberta-large") def __init__(self, path, members=None, backend="auto", device=None, dtype=None, use_gazetteer=True, batch_size=None, providers=None, threads=None): members = list(members or self.MEMBERS) self.backend = self._pick_backend(path, members, backend) if self.backend == "torch": if device is None: device = "cuda" if torch.cuda.is_available() else "cpu" if dtype is None: dtype = torch.bfloat16 if str(device).startswith("cuda") else torch.float32 self.members = [_TorchMember(os.path.join(path, m), device, dtype, batch_size or 64) for m in members] else: self.members = [_OnnxMember(os.path.join(path, m), os.path.join(path, "onnx", m), providers, batch_size or 8, threads) for m in members] if len(members) == len(self.MEMBERS): # The ensemble is calibrated as its own system: each member's bias # describes its own posterior, not the average of two. with open(os.path.join(path, "ensemble_config.json"), encoding="utf-8") as f: ens = json.load(f) self.weights = ens["weights"] self.bias = ens["punct_bias"] self.bias_per_lang = ens["punct_bias_per_lang"] else: self.weights = [1.0 / len(members)] * len(members) meta = self.members[0].meta if len(members) == 1 else {} self.bias = meta.get("punct_bias") self.bias_per_lang = meta.get("punct_bias_per_lang") or {} gz = os.path.join(path, "gazetteer.json") self.gazetteer = {} if use_gazetteer and os.path.exists(gz): with open(gz, encoding="utf-8") as f: self.gazetteer = json.load(f) @staticmethod def _pick_backend(path, members, backend): if backend == "auto": has_st = all(os.path.exists(os.path.join(path, m, "model.safetensors")) for m in members) backend = "torch" if (torch is not None and has_st) else "onnx" if backend == "torch" and torch is None: raise ImportError("backend='torch' needs torch; install it, or use " "backend='onnx' (onnxruntime + tokenizers only)") if backend not in ("torch", "onnx"): raise ValueError("backend must be 'auto', 'torch' or 'onnx'") return backend @classmethod def from_pretrained(cls, repo_or_path="valkayuh/dewpoint", **kw): """A local directory, or a Hugging Face repo id to download. Only the files the chosen backend needs are fetched: the ONNX backend never downloads the safetensors, and the torch backend never downloads ONNX.""" if os.path.isdir(repo_or_path): return cls(repo_or_path, **kw) from huggingface_hub import snapshot_download members = kw.get("members") or cls.MEMBERS backend = kw.get("backend", "auto") if backend == "auto": backend = "torch" if torch is not None else "onnx" kw["backend"] = backend allow = ["dewpoint.py", "ensemble_config.json", "gazetteer.json"] for m in members: allow += [f"{m}/vpunct_config.json", f"{m}/tokenizer.json", f"{m}/tokenizer_config.json"] allow += ([f"{m}/config.json", f"{m}/model.safetensors"] if backend == "torch" else [f"onnx/{m}/*"]) return cls(snapshot_download(repo_or_path, allow_patterns=allow), **kw) # ------------------------------------------------------------ core def _bias_for(self, lang): b = None if self.bias_per_lang: b = self.bias_per_lang.get(lang) or self.bias_per_lang.get(base_lang(lang)) return b if b is not None else self.bias def predict(self, words: Sequence[str], lang: str = "en") -> dict: """Per-word labels and averaged posteriors for an already-split word list. ``words`` should be lowercased and free of punctuation, as ASR output is. """ words = list(words) acc_p = acc_c = None for w, m in zip(self.weights, self.members): p, c = m.posteriors(words, lang) acc_p = w * p if acc_p is None else acc_p + w * p acc_c = w * c if acc_c is None else acc_c + w * c lp = np.log(np.clip(acc_p, 1e-9, None)) b = self._bias_for(lang) if b is not None: lp = lp + np.asarray(b, dtype=np.float32)[None, :] punct = [PUNCT_LABELS[i] for i in lp.argmax(1)] if len(words) else [] case = [CASE_LABELS[i] for i in acc_c.argmax(1)] if len(words) else [] return {"words": words, "punct": punct, "case": case, "punct_probs": acc_p, "case_probs": acc_c, "punct_scores": lp} @staticmethod def _close(punct, scores): """A complete text ends on a sentence mark: if the last word got none, take the likeliest of . ? ! for it. Applied to finished text only -- predict() is left exactly as benchmarked.""" if punct and punct[-1] not in _SENT_END: ends = [PUNCT_LABELS.index(x) for x in ("PERIOD", "QUESTION", "EXCLAM")] punct[-1] = PUNCT_LABELS[max(ends, key=lambda i: scores[-1][i])] return punct def _postprocess(self, words, punct, case, lang, is_start=True, prev_punct=None): """Gazetteer surface forms (iPhone, McDonald) and a capital after every sentence-final mark.""" case = list(case) caseless = is_caseless(lang) gz = self.gazetteer.get(base_lang(lang), {}) out = [] for i, w in enumerate(words): s = gz.get(w) if s and not caseless: out.append(s) case[i] = "AS_IS" else: out.append(w) if not caseless and words: lead = prev_punct in _SENT_END if prev_punct is not None else False if case[0] == "LOWER" and (is_start or lead): case[0] = "CAP" for i in range(1, len(words)): if punct[i - 1] in _SENT_END and case[i] == "LOWER": case[i] = "CAP" return out, list(punct), case def _render(self, words, punct, case, lang): surf = surface_for(lang) sep = "" if is_no_space(lang) else " " return sep.join((w if c == "AS_IS" else apply_case(w, c)) + surf.get(p, "") for w, p, c in zip(words, punct, case)) def restore(self, text: str, lang: str = "en", close: bool = True) -> str: """Unpunctuated text in, punctuated and truecased text out. ``close`` ends the text on a sentence mark; turn it off when ``text`` is a fragment that continues elsewhere. """ words = split_words(text, lang) if not words: return text r = self.predict(words, lang) punct = self._close(r["punct"], r["punct_scores"]) if close else r["punct"] w, p, c = self._postprocess(words, punct, r["case"], lang) return self._render(w, p, c, lang) def restore_batch(self, texts: Iterable[str], lang: str = "en") -> list[str]: return [self.restore(t, lang) for t in texts] def stream(self, lang="en", lag=6, context=180, stable_n=2): return StreamingPunctuator(self, lang, lag, context, stable_n) class StreamingPunctuator: """Commit-with-lag for live transcripts. Later words change earlier decisions -- "what do you think" only becomes a question at its last word -- so a word is released only once it has ``lag`` words of right context and the same label for ``stable_n`` consecutive updates. Call ``push(words)`` as ASR emits, ``finish()`` at the end; each returns the newly committed text. """ def __init__(self, punctuator, lang="en", lag=6, context=180, stable_n=2): self.p, self.lang = punctuator, lang self.lag, self.context, self.stable_n = lag, context, stable_n self.words, self.emitted = [], 0 self._last, self._streak = {}, {} def push(self, new_words) -> str: if isinstance(new_words, str): new_words = split_words(new_words, self.lang) self.words.extend(w.lower() for w in new_words) return self._commit(final=False) def finish(self) -> str: return self._commit(final=True) def _commit(self, final): if not self.words: return "" off = max(0, len(self.words) - self.context) r = self.p.predict(self.words[off:], self.lang) punct, case = r["punct"], r["case"] if final: punct = self.p._close(punct, r["punct_scores"]) limit = len(self.words) if final else len(self.words) - self.lag done, i = [], self.emitted while i < limit: j = i - off if j < 0: i += 1 continue if not final: lab = (punct[j], case[j]) if self._last.get(i) == lab: self._streak[i] = self._streak.get(i, 1) + 1 else: self._last[i], self._streak[i] = lab, 1 if self._streak[i] < self.stable_n: break done.append(i) i += 1 if not done: return "" s, e = done[0], done[-1] + 1 prev = punct[s - off - 1] if s - off - 1 >= 0 else None w, p, c = self.p._postprocess(self.words[s:e], punct[s - off:e - off], case[s - off:e - off], self.lang, is_start=(s == 0), prev_punct=prev) self.emitted = e out = self.p._render(w, p, c, self.lang) return out if s == 0 or is_no_space(self.lang) else " " + out if __name__ == "__main__": import argparse import sys ap = argparse.ArgumentParser(description="Restore punctuation and case.") ap.add_argument("text", nargs="*", help="text to restore (default: stdin)") ap.add_argument("--lang", default="en", help="ISO 639-1 code, e.g. en, de, zh") ap.add_argument("--model", default=os.path.dirname(os.path.abspath(__file__))) ap.add_argument("--single", action="store_true", help="use only the mmBERT-base member (307M, faster)") ap.add_argument("--backend", default="auto", choices=["auto", "torch", "onnx"]) a = ap.parse_args() p = Punctuator.from_pretrained(a.model, backend=a.backend, members=["mmbert-base"] if a.single else None) src = [" ".join(a.text)] if a.text else sys.stdin.read().splitlines() for line in src: print(p.restore(line, a.lang))