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())}
|