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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support