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)