File size: 4,915 Bytes
6cc3500 | 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 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | """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)
|