File size: 5,676 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
131
132
133
134
135
136
137
138
139
140
141
142
143
"""V3 解碼器:lattice 上的 Viterbi + 似然比檢定 + 最小 splice。

決策規則:非 keep 候選必須以 per-rule-group margin τ 勝過 keep
(u = logit[cand] − logit[keep] − τ > 0 才成為選項),
再以 DP 選出總效用最大的**不重疊**編輯集合——重疊的邊在這裡競爭,
取代 V2 的「最長優先預先裁剪」。

Determinism:效用嚴格大於才更新(tie 傾向 keep/先做出的決策),
無隨機性,同輸入必同輸出。

輸出是對 base 文本的最小 splice 編輯清單:非站點區段一個位元組都不動,
從結構上根除 V2 renderer 的間距/標點慣例劣勢。
"""
from __future__ import annotations

import math
from dataclasses import dataclass

import numpy as np

from twlat.lattice import TAU_GROUP, Lattice, LatticeBuilder

DEFAULT_TAU = {"variant": 0.0, "lexical": 0.0, "style": 0.0}

# 已驗證的操作點(neutral-dev 校準,數字見技術報告 §18.16)。
# 三者的差別只在「要多少證據才動手」,模型與字典完全相同。
PRESETS = {
    # 最大化 benchmark site accuracy:要求 20:1 勝算才改動。
    # 副作用:孤立短句中 網絡→網路(8:1)這類正確改動會被壓掉。
    "accuracy": {"variant": 0.0, "lexical": 3.0, "style": 6.0},
    # 產品預設:2.7:1 勝算即改動。主觀行為符合直覺,benchmark 代價 −0.4pp。
    "balanced": {"variant": 0.0, "lexical": 1.0, "style": 3.0},
    # 最大召回:模型認為較可能就改(僅硬過濾與 input_only 把關)。
    "aggressive": {"variant": 0.0, "lexical": 0.0, "style": 0.0},
}

# fo_bonus 建議值(配合 PRESETS 使用)。陸式專用形式(服務器/網絡/軟件/視頻,
# 見 LatticeBuilder.cn_only)保留時扣分——字典說它們不該是輸出。
# benchmark 代價 −0.28pp(gold 本身含這些形式,見報告 §18.17)。
FO_BONUS = {"accuracy": 0.0, "balanced": 0.0, "aggressive": 0.0,
            "taiwanize": 4.0}
PRESETS["taiwanize"] = dict(PRESETS["balanced"])


@dataclass
class Edit:
    start: int
    end: int
    replacement: str
    observed: str
    utility: float
    rule_type: str


def decode(lb: LatticeBuilder, lat: Lattice, logits: np.ndarray,
           tau: dict[str, float] | None = None,
           fo_bonus: float = 0.0) -> list[Edit]:
    """logits: [n_edges, C],與 lat.edges 對齊(C 為該批的候選欄數)。

    τ 的量綱:候選與 keep 的 logit 差**就是**模型 softmax 下的對數機率比
    (log-softmax 對每列減去同一常數,差不變),因此 τ 可直接讀成勝算比門檻——
    τ=1 ≈ 2.7:1、τ=3 ≈ 20:1。PRESETS 提供三個已驗證的操作點。
    """
    tau = tau or DEFAULT_TAU
    options: list[tuple[int, int, float, str, str, str]] = []
    for i, e in enumerate(lat.edges):
        if i >= len(logits):
            break
        g = lb.groups[e.gid]
        members = [lb.strings[x] for x in g["m"]]
        obs = members[e.obs_ix]
        keep_s = float(logits[i, e.obs_ix])
        if not math.isfinite(keep_s):
            continue
        # from_only 先驗:字典明確不背書 observed 作為輸出(服務器/視頻/博客)。
        # 這類形式在 C3 網爬語料中大量出現且被標為 keep(實測 1% 資料中
        # 服務器 有 12 筆 keep、0 筆 change),模型因此學到保留。
        # 字典知識在解碼層補回:保留它需要額外證據。
        if fo_bonus and lb.cn_only.get(e.gid, [False] * len(members))[e.obs_ix]:
            keep_s -= fo_bonus
        for j, cand in enumerate(members):
            if j == e.obs_ix or j >= logits.shape[1]:
                continue
            if j < len(e.cand_kill) and e.cand_kill[j]:
                continue
            if g["io"][j]:                      # input_only 成員不可被引入
                continue
            if g.get("fo", [False] * len(members))[j]:
                continue                        # 無規則背書為輸出(視頻/軟件)
            s = float(logits[i, j])
            if not math.isfinite(s):
                continue
            rid = lb.pairs.get((obs, cand), g["r"][j])
            grp = TAU_GROUP.get(lb.rules[rid]["t"], "lexical")
            u = s - keep_s - tau.get(grp, 0.0)
            if u > 1e-9:
                options.append((e.start, e.end, u, cand, obs,
                                lb.rules[rid]["t"]))

    if not options:
        return []

    n = len(lat.text)
    best = np.zeros(n + 1)
    back: list[tuple | None] = [None] * (n + 1)
    by_end: dict[int, list] = {}
    for o in options:
        by_end.setdefault(o[1], []).append(o)
    for opts in by_end.values():
        opts.sort(key=lambda o: (o[0], -o[2]))   # 固定順序 → determinism

    for p in range(1, n + 1):
        best[p] = best[p - 1]
        back[p] = None
        for o in by_end.get(p, []):
            cand_score = best[o[0]] + o[2]
            if cand_score > best[p] + 1e-9:      # 嚴格大於:tie 傾向 keep
                best[p] = cand_score
                back[p] = o

    edits: list[Edit] = []
    p = n
    while p > 0:
        o = back[p]
        if o is None:
            p -= 1
        else:
            edits.append(Edit(o[0], o[1], o[3], o[4], o[2], o[5]))
            p = o[0]
    edits.reverse()
    return edits


def splice(text: str, edits: list[Edit]) -> str:
    """最小編輯:只替換編輯 span,其餘位元組原樣。"""
    out, prev = [], 0
    for e in edits:
        out.append(text[prev:e.start])
        out.append(e.replacement)
        prev = e.end
    out.append(text[prev:])
    return "".join(out)