File size: 15,016 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
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
"""Conversion lattice:V3 的核心資料結構。

對 safe_normalize 後的文本,用 lattice lexicon(dict/lattice_lexicon.json)
建出「所有字典允許的改寫」構成的圖:

    節點 = 字元位置
    邊   = (span, confusion group),group 的每個成員是一個候選(含 keep)

字典知識的分工(PI 指示的問題拆解):
  - **確定性可判的,lattice 直接判**:exceptions 例外詞(函式庫 內不得改 函式)、
    positional_clues(好|消息 不觸發 消息→訊息)、詞界穿越(商調|制度 的 調製 邊)
    ——這些命中即砍邊/砍候選,模型看不到也不需要看。
  - **語境相依的,交給模型**:剩下的每條邊帶 64 維字典特徵
    (領域 one-hot、規則型別、正反 clue 命中、語料頻率、editorial confidence…),
    模型只回答「這個語境下哪個成員成立」。

重疊的邊一律保留,交給 Viterbi 全域解碼(src/twlat/decoder.py)。
"""
from __future__ import annotations

import json
import math
import pathlib
from dataclasses import dataclass, field

import ahocorasick
import numpy as np

from twlat.paths import data_file

LEXICON_PATH = data_file("dict/lattice_lexicon.json")

# ---- 特徵配置(改動任何索引都要 bump lexicon SCHEMA_VERSION)----
N_DOMAINS = 36          # 33 實際領域 + other + 無標記 + 保留
RULE_TYPES = ["cross_strait", "variant_char", "tw_phrase", "confusable",
              "ai_filler", "translationese", "variant", "political_coloring",
              "typo", "other"]
FEAT_DIM = 64
CLUE_WINDOW = 40        # clue 比對視窗(±40 字),與 V2 preprocess_r 一致
CN_ONLY_RATIO = 0.35    # 陸式專用形式的頻率比上限(見 LatticeBuilder.cn_only)

# τ 校準用的規則分組(decoder 對非 keep 邊套 per-group margin)
TAU_GROUP = {"variant_char": "variant", "variant": "variant",
             "cross_strait": "lexical", "confusable": "lexical",
             "tw_phrase": "lexical", "typo": "lexical",
             "ai_filler": "style", "translationese": "style",
             "political_coloring": "style"}


@dataclass
class Edge:
    start: int
    end: int
    gid: int                      # confusion group id
    obs_ix: int                   # observed 形式在 group 正規順序中的 index
    cand_kill: list[bool]         # 各成員是否被硬過濾砍除(observed 永不砍)
    word_crossing: bool = False   # 邊穿越詞界(軟特徵;variant_char 穿越則硬砍)
    word_contained: bool = False  # 邊嚴格位於某個已知詞內(軟特徵)


@dataclass
class Lattice:
    text: str
    edges: list[Edge] = field(default_factory=list)


class LatticeBuilder:
    def __init__(self, lexicon_path: pathlib.Path = LEXICON_PATH):
        lex = json.loads(pathlib.Path(lexicon_path).read_text(encoding="utf-8"))
        self.version: str = lex["version"]
        self.strings: list[str] = lex["strings"]
        self.groups: list[dict] = lex["groups"]
        self.form2group: dict[str, int] = lex["form2group"]
        self.rules: list[dict] = lex["rules"]
        self.pairs: dict[tuple[str, str], int] = {
            tuple(k.split("\t")): v for k, v in lex["pairs"].items()}
        self.freq: dict[str, int] = lex["freq"]
        self.exceptions: dict[str, list[int]] = lex["exceptions"]

        self._a_sites = ahocorasick.Automaton()
        for f in self.form2group:
            self._a_sites.add_word(f, f)
        self._a_sites.make_automaton()

        self._a_exc = ahocorasick.Automaton()
        for s in self.exceptions:
            self._a_exc.add_word(s, s)
        self._a_exc.make_automaton()

        self._a_words = ahocorasick.Automaton()
        for w in lex["word_forms"]:
            self._a_words.add_word(w, w)
        self._a_words.make_automaton()

        # 陸式專用詞形:字典不背書為輸出(fo)**且**臺灣語料頻率遠低於
        # 組內最高(< CN_ONLY_RATIO)。單靠 fo 不夠精確——項目/提升/設備
        # 也只出現在某些規則的 from 側,但它們是正常臺灣詞(頻率比 ≈ 1.0)。
        # 服務器 0.115、網絡 0.151、軟件 0.190、視頻 0.295 才是真正的陸式形式。
        self.cn_only: dict[int, list[bool]] = {}
        for gid, g in enumerate(self.groups):
            mem = [self.strings[i] for i in g["m"]]
            top = max(self.freq.get(m, 0) for m in mem) or 1
            self.cn_only[gid] = [
                bool(fo) and self.freq.get(m, 0) / top < CN_ONLY_RATIO
                for m, fo in zip(mem, g.get("fo", [False] * len(mem)))]

    # ---- 建圖 ----

    def build(self, text: str) -> Lattice:
        lat = Lattice(text=text)

        # 例外詞 span(規則相依):exception 覆蓋邊 → 砍該規則的非 keep 候選
        exc_spans: list[tuple[int, int, list[int]]] = []
        for end, s in self._a_exc.iter(text):
            exc_spans.append((end - len(s) + 1, end + 1, self.exceptions[s]))

        # 詞界證據:最長優先不重疊
        word_spans = self._longest_nonoverlap(self._a_words.iter(text))

        for end, form in self._a_sites.iter(text):
            start = end - len(form) + 1
            end = end + 1
            gid = self.form2group[form]
            g = self.groups[gid]
            members = [self.strings[i] for i in g["m"]]
            obs_ix = members.index(form)

            crossing, contained = self._word_relation(start, end, word_spans)
            g_type = self.rules[g["r"][obs_ix]]["t"]
            if crossing and g_type == "variant_char":
                continue   # 字級變體穿越詞界(商調|制度 的 調製)→ 整條邊砍掉

            kill = [False] * len(members)
            for ci, cand in enumerate(members):
                if ci == obs_ix:
                    continue
                rid = self.pairs.get((form, cand), g["r"][ci])
                rule = self.rules[rid]
                if self._exception_hit(start, end, rid, exc_spans):
                    kill[ci] = True
                elif not self._positional_ok(text, start, end, rule):
                    kill[ci] = True
            lat.edges.append(Edge(start, end, gid, obs_ix, kill,
                                  crossing, contained))
        lat.edges.sort(key=lambda e: (e.start, -(e.end - e.start), e.gid))
        return lat

    @staticmethod
    def _longest_nonoverlap(hits) -> list[tuple[int, int]]:
        spans = sorted(((end - len(w) + 1, end + 1) for end, w in hits),
                       key=lambda s: (s[0], -(s[1] - s[0])))
        out: list[tuple[int, int]] = []
        last = -1
        for s, e in spans:
            if s >= last:
                out.append((s, e))
                last = e
        return out

    @staticmethod
    def _word_relation(start: int, end: int,
                       word_spans: list[tuple[int, int]]) -> tuple[bool, bool]:
        crossing = contained = False
        for ws, we in word_spans:
            if we <= start:
                continue
            if ws >= end:
                break
            if (ws < start < we < end) or (start < ws < end < we):
                crossing = True
            if ws <= start and end <= we and (ws, we) != (start, end):
                contained = True
        return crossing, contained

    @staticmethod
    def _exception_hit(start: int, end: int, rid: int,
                       exc_spans: list[tuple[int, int, list[int]]]) -> bool:
        for xs, xe, rids in exc_spans:
            if xs <= start and end <= xe and (xs, xe) != (start, end) \
                    and rid in rids:
                return True
        return False

    @staticmethod
    def _positional_ok(text: str, start: int, end: int, rule: dict) -> bool:
        """負向 positional 是否決;正向(before/after)存在時須至少滿足一個。"""
        pos_req, pos_ok = False, False
        for kind, arg in rule["po"]:
            if kind == "not_after" and text[max(0, start - len(arg)):start] == arg:
                return False
            if kind == "not_before" and text[end:end + len(arg)] == arg:
                return False
            if kind in ("before", "after"):
                pos_req = True
                if kind == "before" and text[end:end + len(arg)] == arg:
                    pos_ok = True
                if kind == "after" and text[max(0, start - len(arg)):start] == arg:
                    pos_ok = True
        return pos_ok if pos_req else True

    # ---- 特徵 ----

    def edge_features(self, lat: Lattice, edge: Edge,
                      reveal_observed: bool) -> np.ndarray:
        """[C, FEAT_DIM]。reveal_observed=False 用於 cloze 預訓練:
        observed 相依的維度(is_keep、長度差)歸零,避免標籤洩漏。"""
        g = self.groups[edge.gid]
        members = [self.strings[i] for i in g["m"]]
        obs = members[edge.obs_ix]
        lo = max(0, edge.start - CLUE_WINDOW)
        ctx = lat.text[lo:edge.end + CLUE_WINDOW]
        top_freq = max(self.freq.get(m, 0) for m in members)

        out = np.zeros((len(members), FEAT_DIM), dtype=np.float32)
        for ci, cand in enumerate(members):
            rid = self.pairs.get((obs, cand), g["r"][ci])
            rule = self.rules[rid]
            f = out[ci]
            if reveal_observed:
                f[0] = float(ci == edge.obs_ix)
                f[55] = (len(cand) - len(obs)) / 6.0
            t_ix = RULE_TYPES.index(rule["t"]) if rule["t"] in RULE_TYPES \
                else RULE_TYPES.index("other")
            f[1 + t_ix] = 1.0
            for d in rule["d"]:
                if d < N_DOMAINS - 3:
                    f[11 + d] = 1.0
            if not rule["d"]:
                f[11 + N_DOMAINS - 2] = 1.0     # 無領域標記
            f[47] = min(sum(1 for c in rule["pc"] if c in ctx), 5) / 5.0
            f[48] = min(sum(1 for c in rule["nc"] if c in ctx), 5) / 5.0
            f[49] = float(bool(rule["en"]) and rule["en"].lower()
                          in lat.text.lower())
            fq = self.freq.get(cand, 0)
            f[50] = math.log10(fq + 1) / 7.0
            f[51] = {None: 0.5, "low": 0.0, "high": 1.0}.get(rule["cf"], 0.5)
            f[52] = float(fq == top_freq)
            f[53] = len(cand) / 6.0
            f[54] = len(members) / 8.0
            f[56] = float(edge.word_contained)
            f[57] = float(edge.word_crossing)
            f[58] = float(g["io"][ci])
            f[59] = float(edge.cand_kill[ci])
        return out

    # ---- 序列化(訓練 shard 用)----

    def to_arrays(self, lat: Lattice, s_max: int, c_max: int
                  ) -> dict[str, np.ndarray] | None:
        """定長陣列。site_gold = obs_ix(真實語料上 observed 即正解)。
        溢出時依優先序裁邊:lexical 規則邊 > 可遮罩 variant > 不可遮罩 variant。"""
        edges = lat.edges
        if len(edges) > s_max:
            def prio(e: Edge):
                g = self.groups[e.gid]
                t = self.rules[g["r"][e.obs_ix]]["t"]
                return (0 if TAU_GROUP.get(t) != "variant" else
                        1 if g["mk"][e.obs_ix] else 2)
            edges = sorted(edges, key=lambda e: (prio(e), e.start))[:s_max]
            edges.sort(key=lambda e: (e.start, -(e.end - e.start), e.gid))

        n = len(edges)
        if n == 0:
            return None
        arr = {
            "site_span": np.zeros((s_max, 2), dtype=np.int16),
            "site_gid": np.full(s_max, -1, dtype=np.int32),
            "site_gold": np.zeros(s_max, dtype=np.int8),
            "site_ncand": np.zeros(s_max, dtype=np.int8),
            "site_kill": np.zeros((s_max, c_max), dtype=bool),
            "site_maskable": np.zeros(s_max, dtype=bool),
            "n_sites": np.int16(n),
        }
        for i, e in enumerate(edges):
            g = self.groups[e.gid]
            arr["site_span"][i] = (e.start, e.end)
            arr["site_gid"][i] = e.gid
            arr["site_gold"][i] = e.obs_ix
            arr["site_ncand"][i] = min(len(g["m"]), c_max)
            arr["site_kill"][i, :len(e.cand_kill)] = e.cand_kill[:c_max]
            arr["site_maskable"][i] = g["mk"][e.obs_ix]
        return arr


MASK_ID = 2


def build_masked_view(ids: np.ndarray, spans: np.ndarray, maskable: np.ndarray,
                      mask_id: int = MASK_ID
                      ) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    """把可遮罩站點收合成單一 MASK token 的序列視圖。

    訓練 collate 與 runtime **必須共用此函式**——遮罩政策只依 observed 形式
    (maskable 是 form 的函數),兩邊分佈才一致。

    重疊站點共用 MASK:以 (start, -len) 貪婪選出不重疊的遮罩單元,
    其餘站點的 m_span 透過索引投影落在覆蓋它的 MASK 位置(可含殘餘可見字元)。

    回傳 (masked_ids, m_spans[n,2], old2new[len+1])。
    """
    n = len(ids)
    order = sorted(range(len(spans)),
                   key=lambda i: (int(spans[i][0]), -(int(spans[i][1]) - int(spans[i][0]))))
    units: list[tuple[int, int]] = []
    last = -1
    for i in order:
        if not maskable[i]:
            continue
        s, e = int(spans[i][0]), int(spans[i][1])
        if s >= last:
            units.append((s, e))
            last = e

    old2new = np.zeros(n + 1, np.int32)
    segs: list[np.ndarray] = []
    prev = pos = 0
    mask_tok = np.array([mask_id], dtype=ids.dtype)
    for s, e in units:
        for k in range(prev, s):
            old2new[k] = pos + (k - prev)
        pos += s - prev
        segs.append(ids[prev:s])
        segs.append(mask_tok)
        for k in range(s, e):
            old2new[k] = pos
        pos += 1
        prev = e
    for k in range(prev, n):
        old2new[k] = pos + (k - prev)
    pos += n - prev
    segs.append(ids[prev:n])
    old2new[n] = pos
    masked = np.concatenate(segs) if segs else ids[:0]

    m_spans = np.zeros((len(spans), 2), np.int32)
    for i, (s, e) in enumerate(spans):
        a = int(old2new[int(s)])
        b = max(int(old2new[int(e)]), a + 1)
        m_spans[i] = (a, b)
    return masked, m_spans, old2new


def density_report(builder: LatticeBuilder, texts: list[str]) -> dict:
    """邊密度統計,決定 S_MAX。"""
    import collections
    ns, per_type = [], collections.Counter()
    for t in texts:
        lat = builder.build(t)
        ns.append(len(lat.edges))
        for e in lat.edges:
            g = builder.groups[e.gid]
            per_type[builder.rules[g["r"][e.obs_ix]]["t"]] += 1
    ns_arr = np.array(ns)
    return {"n_texts": len(texts),
            "sites_mean": round(float(ns_arr.mean()), 2),
            "sites_p50": int(np.percentile(ns_arr, 50)),
            "sites_p95": int(np.percentile(ns_arr, 95)),
            "sites_p99": int(np.percentile(ns_arr, 99)),
            "sites_max": int(ns_arr.max()),
            "per_type": dict(per_type.most_common())}