Twinity-1 / twlat /__init__.py
JacobLinCool's picture
Twinity-1: weights, compiled dictionary, inference code
6cc3500 verified
Raw
History Blame Contribute Delete
4.92 kB
"""TWLAT — 中國大陸中文 → 臺灣正體中文的確定性轉換器。
以字典編譯的 conversion lattice 界定動作空間,8.9M 參數的模型只裁決
語境相依的歧義,Viterbi 全域解碼 + 最小 splice 產生輸出。
>>> import twlat
>>> twlat.convert("这个程序有bug,请在服务器上重新部署。")
'這個程式有 bug,請在伺服器上重新部署。'
>>> conv = twlat.Converter(device="cpu")
>>> conv.convert_batch(["文本一", "文本二"])
['文本一', '文本二']
熱更新:`dict/lattice_lexicon.json` 重新編譯即可新增字典條目,
不需要重新訓練(候選由共用 char embedding 動態編碼)。
"""
from __future__ import annotations
import functools
import pathlib
__version__ = "0.3.1"
__all__ = ["Converter", "convert", "convert_batch", "Decision", "Result",
"DEFAULT_CKPT", "__version__"]
REPO = pathlib.Path(__file__).resolve().parents[2]
DEFAULT_CKPT = REPO / "runs/v3/best.pt"
class Decision:
"""模型對單一 lattice 邊做出的改寫決策。"""
__slots__ = ("start", "end", "source", "target", "utility", "rule_type")
def __init__(self, d: dict):
self.start, self.end = d["span"]
self.source = d["from"]
self.target = d["to"]
self.utility = d["utility"]
self.rule_type = d["rule_type"]
def __repr__(self) -> str:
return (f"Decision({self.source!r}{self.target!r} @[{self.start},"
f"{self.end}) {self.rule_type} u={self.utility:.2f})")
class Result:
"""單一段落的轉換結果。`str(result)` 即輸出文字。"""
__slots__ = ("text", "decisions", "fast_path")
def __init__(self, r):
self.text: str = r.output
self.decisions: list[Decision] = [Decision(d) for d in r.decisions]
self.fast_path: bool = r.fast_path
def __str__(self) -> str:
return self.text
def __repr__(self) -> str:
return f"Result({self.text!r}, {len(self.decisions)} decisions)"
class Converter:
"""可重複使用的轉換器(載入一次模型,之後重複呼叫)。
Parameters
----------
ckpt : 模型 checkpoint 路徑(預設 runs/v3/best.pt)
device : "cpu" / "mps" / "cuda";預設自動偵測。CPU 單執行緒吞吐最佳
(~11k 字/秒),見技術報告 §18.15。
preset : 操作點(見 twlat.decoder.PRESETS)——
"accuracy"(benchmark 最佳,最保守)、
"balanced"(產品預設)、"taiwanize"(額外修正陸式專用詞)、
"aggressive"(最大召回)。傳 None 則用 model/tau_v3.json。
tau : 直接指定 per-rule-group 門檻(覆寫 preset 的 τ)。
lexicon : 替代 lattice lexicon(熱更新用)。
"""
def __init__(self, ckpt=None, device: str | None = None, tau=None,
lexicon=None, preset: str | None = "balanced"):
import sys
sys.path[:0] = [str(REPO / "src")]
from twlat.decoder import FO_BONUS, PRESETS
from twlat.runtime_v3 import TWLATV3Runtime
fo = 0.0
if preset is not None:
if preset not in PRESETS:
raise ValueError(f"preset 須為 {sorted(PRESETS)}")
tau = tau or PRESETS[preset]
fo = FO_BONUS.get(preset, 0.0)
self.preset = preset
self._rt = TWLATV3Runtime(str(ckpt or DEFAULT_CKPT), device=device,
tau=tau, lexicon_path=lexicon, fo_bonus=fo)
def convert(self, text: str) -> str:
"""單段轉換,回傳文字。"""
return self._rt.convert_batch([text])[0].output
def convert_batch(self, texts: list[str], batch_size: int = 8) -> list[str]:
"""批次轉換。batch_size=8 為 CPU 最佳操作點。"""
return [r.output for r in self._rt.convert_batch(texts, batch_size)]
def explain(self, text: str) -> Result:
"""回傳含逐項決策的結果(span、來源形式、目標形式、效用、規則型別)。"""
return Result(self._rt.convert_batch([text])[0])
def explain_batch(self, texts: list[str], batch_size: int = 8) -> list[Result]:
return [Result(r) for r in self._rt.convert_batch(texts, batch_size)]
@property
def lexicon_version(self) -> str:
return self._rt.lb.version
@functools.lru_cache(maxsize=2)
def _default(device: str | None = None) -> Converter:
return Converter(device=device)
def convert(text: str, device: str | None = None) -> str:
"""便利函式:轉換單段文字(首次呼叫會載入模型並快取)。"""
return _default(device).convert(text)
def convert_batch(texts: list[str], device: str | None = None,
batch_size: int = 8) -> list[str]:
"""便利函式:批次轉換。"""
return _default(device).convert_batch(texts, batch_size)