Twinity-1 / twlat /lattice.py
JacobLinCool's picture
Twinity-1: weights, compiled dictionary, inference code
6cc3500 verified
Raw
History Blame Contribute Delete
15 kB
"""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())}