0c D512 Punctuation Restoration (Bidirectional Encoder)
Project website: 0c.biggo.com
Restores punctuation in unpunctuated Traditional Chinese text. The model inserts punctuation while preserving all source characters. It does not append punctuation at the end of the text. Labels: οΌ γ οΌ οΌ γ οΌ, plus a "no insertion" class.
- Architecture: a bidirectional Transformer encoder with 8 layers, hidden size 512, 8 attention heads, FFN size 768. A D512 backbone pretrained from random initialization was adapted to bidirectional attention and fine-tuned for punctuation restoration.
- Window: 128 characters; stride: 64
- Format: ONNX (opset 18), CPU inference
- Vocabulary: 27,857 characters (
vocab.json)
Files
| File | Description |
|---|---|
model.onnx |
Punctuation classification model; SHA-256 4db3ed40df777347d9c455da01285b01de192d4c210c5f7cd9f0646c727880d8 |
vocab.json |
Special tokens and character vocabulary; SHA-256 556dffaffbf4c05211a846dcee55ef735cea68c5754f9bffcc6f9588124172b3 |
manifest.json |
Labels, threshold, window settings, and file hashes |
Evaluation
Export validation used 1,171 development segments that had been seen during training, rather than an unseen test set. At threshold 0.63, precision is 0.897, recall is 0.755, and F1 is 0.820. ONNX and PyTorch produce identical punctuation insertions for every validation segment.
Installation
Requires Python 3.9 or later. The examples use CPU inference.
pip install onnxruntime numpy huggingface_hub
The first run downloads the model from Hugging Face (approximately 130 MB). Subsequent runs use the local cache.
Usage
Save the complete example below as punctuate.py and run it. Threshold 0.63 is an operating point
with precision around 0.9 in the reported validation. Lower thresholds insert more punctuation but also increase incorrect insertions.
import json
import re
import numpy as np
import onnxruntime as ort
from huggingface_hub import snapshot_download
path = snapshot_download("Funmula/0c-punct-d512")
manifest = json.load(open(f"{path}/manifest.json", encoding="utf-8"))
vocab = json.load(open(f"{path}/vocab.json", encoding="utf-8"))
CHARS = {c: i + 8 for i, c in enumerate(vocab["characters"])} # IDs 0β7 are special tokens; 4 = UNK
LABELS = manifest["labels"] # Class 0 = no insertion
THRESHOLD, WINDOW, STRIDE = manifest["threshold"], manifest["window"], manifest["stride"]
session = ort.InferenceSession(f"{path}/model.onnx", providers=["CPUExecutionProvider"])
URL = re.compile(r"(?:https?://|www\.)[^\sοΌγοΌοΌγοΌγγγγοΌοΌγγ]+", re.I)
ENUMERATORS = set("δΈδΊδΈεδΊε
δΈε
«δΉεε£Ήθ²³εθδΌιΈζζηζΎη²δΉδΈδΈζε·±εΊθΎε£¬ηΈ")
def is_han(c):
return "γ" <= c <= "ιΏΏ" or "\U00020000" <= c <= "\U000323af"
def allowed(text):
# Insert after a Han character outside URLs, followed by Han text or whitespace;
# the next non-whitespace character must be Han or ASCII alphanumeric.
protected = {i for m in URL.finditer(text) for i in range(m.start(), m.end())}
out = []
for i, c in enumerate(text):
nxt = text[i + 1:i + 2]
visible = text[i + 1:].lstrip()[:1]
out.append(is_han(c) and i not in protected
and (not nxt or is_han(nxt) or nxt.isspace())
and (not visible or is_han(visible) or visible.isascii() and visible.isalnum()))
return out
def window_probs(text):
ids = np.array([[1, 7] + [CHARS.get(c, 4) for c in text]], dtype=np.int64) # [BOS][TASK] source text
return session.run(["probs"], {"ids": ids})[0][2:] # Row i: label probabilities after character i
def probabilities(text):
# For long text, use sliding windows and keep the most central prediction at each position.
n = len(text)
starts = [0] if n <= WINDOW else list(range(0, n - WINDOW, STRIDE)) + [n - WINDOW]
out, best = [None] * n, [-1] * n
for s in starts:
probs = window_probs(text[s:s + WINDOW])
for j in range(len(probs)):
margin = min(j, len(probs) - 1 - j) if n > WINDOW else 0
if margin > best[s + j]:
best[s + j], out[s + j] = margin, probs[j]
return np.stack(out)
def punctuate(text):
# Insert punctuation while preserving every source character; leave the final boundary unchanged.
if not text:
return text
ok, probs = allowed(text), probabilities(text)
out = []
for i, c in enumerate(text):
out.append(c)
if i == len(text) - 1 or not ok[i]:
continue
k = int(probs[i, 1:].argmax()) + 1
if i == 0 and LABELS[k] == "γ" and c in ENUMERATORS: # Skip a list comma after a single initial enumerator
continue
if probs[i, k] >= THRESHOLD:
out.append(LABELS[k])
return "".join(out)
print(punctuate("ε₯½ηζη₯ιδΊζ倩θ¦")) # ε₯½ηζη₯ιδΊοΌζ倩θ¦
The input-method integration processes at most 256 characters per request. The example itself has no length limit and uses sliding windows for longer text.
License and Release Scope
Model weights and the accompanying inference examples are released under the MIT license (see LICENSE).
This release covers this D512 checkpoint only. Training code and training materials are not included.