File size: 7,014 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 | """V3 特徵組裝:pretrain_data(離線)、train_v3(collate)、runtime_v3(線上)
三方共用的唯一實作——訓練與推論的特徵分佈必須 bit-consistent。
分工備忘:
- 靜態特徵(rule type/domain/freq/conf…)以**成員歸屬規則**(group["r"][ci])
編碼進 LexTables.static;
- 語境相依特徵(clue 命中/english anchor)以 **(observed, cand) pair 規則**
在文本上計算(site_arrays)。
兩者的規則來源不同是刻意的:pair 規則才知道「這個轉換方向」的語意條件。
"""
from __future__ import annotations
import json
import math
import pathlib
import numpy as np
import regex
from twlat.paths import data_file
SEQ, S_MAX, C_MAX, L_MAX = 512, 128, 8, 8
HAN_VOCAB = 4096
HASH_SPACE = 59000
FEAT_DIM = 64
CLUE_WINDOW = 40
MASK_ID = 2
HAN = regex.compile(r"\p{Han}")
LATIN = regex.compile(r"[A-Za-z]")
DIGIT = regex.compile(r"\p{Nd}")
PROTECT = regex.compile(r"https?://\S+|[\w.+-]+@[\w-]+\.[\w.]+|`[^`]+`"
r"|[A-Za-z][A-Za-z0-9_.+-]{2,}")
RULE_TYPES = ["cross_strait", "variant_char", "tw_phrase", "confusable",
"ai_filler", "translationese", "variant", "political_coloring",
"typo", "other"]
RT_IX = {t: i for i, t in enumerate(RULE_TYPES)}
def enc_char(ch: str, vocab: dict) -> int:
i = vocab.get(ch)
return i if i is not None else HAN_VOCAB + (ord(ch) % HASH_SPACE)
def text_arrays(text: str, vocab: dict) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""→ (ids int64[n], script uint8[n], prot bool[n])"""
n = len(text)
ids = np.zeros(n, np.int64)
script = np.zeros(n, np.uint8)
prot = np.zeros(n, bool)
for i, ch in enumerate(text):
ids[i] = enc_char(ch, vocab)
script[i] = 1 if HAN.match(ch) else 2 if LATIN.match(ch) else \
3 if DIGIT.match(ch) else 0
for m in PROTECT.finditer(text):
prot[m.start():m.end()] = True
return ids, script, prot
def site_arrays(lb, edges, text: str) -> dict[str, np.ndarray]:
"""lattice edges → 站點中繼陣列(無 gold;gold 由呼叫端投影)。"""
ns = len(edges)
lowered = text.lower()
a = {"span": np.zeros((ns, 2), np.int64),
"gid": np.zeros(ns, np.int32),
"obs": np.zeros(ns, np.int64),
"maskable": np.zeros(ns, bool),
"kill": np.zeros((ns, C_MAX), bool),
"clue": np.zeros((ns, C_MAX, 2), np.uint8),
"eng": np.zeros((ns, C_MAX), bool),
"flags": np.zeros(ns, np.uint8)}
for k, e in enumerate(edges):
g = lb.groups[e.gid]
members = [lb.strings[i] for i in g["m"]]
obs = members[e.obs_ix]
a["span"][k] = (e.start, e.end)
a["gid"][k] = e.gid
a["obs"][k] = e.obs_ix
a["maskable"][k] = g["mk"][e.obs_ix]
a["kill"][k, :len(e.cand_kill)] = e.cand_kill[:C_MAX]
a["flags"][k] = int(e.word_contained) | (int(e.word_crossing) << 1)
ctx = text[max(0, e.start - CLUE_WINDOW):e.end + CLUE_WINDOW]
for ci, cand in enumerate(members[:C_MAX]):
rid = lb.pairs.get((obs, cand), g["r"][ci])
rule = lb.rules[rid]
if rule["pc"]:
a["clue"][k, ci, 0] = min(sum(1 for c in rule["pc"] if c in ctx), 5)
if rule["nc"]:
a["clue"][k, ci, 1] = min(sum(1 for c in rule["nc"] if c in ctx), 5)
if rule["en"]:
a["eng"][k, ci] = rule["en"].lower() in lowered
return a
class LexTables:
"""gid → 候選 token / 靜態特徵 展開表(collate 與 runtime 共用)。"""
def __init__(self, lexicon_path=None, vocab_path=None):
lexicon_path = lexicon_path or data_file("dict/lattice_lexicon.json")
vocab_path = vocab_path or data_file("dict/char_vocab_v3.json")
lex = json.loads(pathlib.Path(lexicon_path).read_text(encoding="utf-8"))
vocab = json.loads(pathlib.Path(vocab_path).read_text(encoding="utf-8"))
self.version = lex["version"]
strings, rules, freq = lex["strings"], lex["rules"], lex["freq"]
G = len(lex["groups"])
self.tok = np.zeros((G, C_MAX, L_MAX), np.int64)
self.ncand = np.zeros(G, np.int8)
self.length = np.zeros((G, C_MAX), np.float32)
self.static = np.zeros((G, C_MAX, FEAT_DIM), np.float32)
self.fo = np.zeros((G, C_MAX), bool)
for gid, g in enumerate(lex["groups"]):
mem = [strings[i] for i in g["m"]][:C_MAX]
for ci, flag in enumerate(g.get("fo", [])[:C_MAX]):
self.fo[gid, ci] = flag
self.ncand[gid] = len(mem)
top = max(freq.get(m, 0) for m in mem)
for ci, m in enumerate(mem):
for k, ch in enumerate(m[:L_MAX]):
self.tok[gid, ci, k] = enc_char(ch, vocab)
self.length[gid, ci] = len(m)
r = rules[g["r"][ci]]
f = self.static[gid, ci]
f[1 + RT_IX.get(r["t"], RT_IX["other"])] = 1.0
for d in r["d"]:
if d < 33:
f[11 + d] = 1.0
if not r["d"]:
f[11 + 34] = 1.0
fq = freq.get(m, 0)
f[50] = math.log10(fq + 1) / 7.0
f[51] = {None: 0.5, "low": 0.0, "high": 1.0}.get(r["cf"], 0.5)
f[52] = float(fq == top)
f[53] = len(m) / 6.0
f[54] = len(mem) / 8.0
f[58] = float(g["io"][ci])
def assemble_cands(lex: LexTables, gid, obs, clue, eng, flags, kill,
reveal_observed: bool):
"""→ (cand_tok, cand_mask, cand_kill, cand_feat),C 裁到本組最大候選數。"""
C = int(lex.ncand[gid].max()) if len(gid) else 1
cand_tok = lex.tok[gid][:, :C]
cand_feat = lex.static[gid][:, :C].copy()
cand_mask = np.arange(C)[None, :] < lex.ncand[gid][:, None]
cand_kill = kill[:, :C].copy()
cand_kill[~cand_mask] = False
cand_feat[:, :, 47] = clue[:, :C, 0] / 5.0
cand_feat[:, :, 48] = clue[:, :C, 1] / 5.0
cand_feat[:, :, 49] = eng[:, :C]
cand_feat[:, :, 56] = (flags & 1)[:, None]
cand_feat[:, :, 57] = ((flags >> 1) & 1)[:, None]
cand_feat[:, :, 59] = cand_kill
cand_feat[:, :, 60] = lex.fo[gid][:, :C]
if reveal_observed:
ar = np.arange(C)[None, :]
cand_feat[:, :, 0] = (ar == obs[:, None]).astype(np.float32)
obs_len = lex.length[gid, obs]
cand_feat[:, :, 55] = (lex.length[gid][:, :C] - obs_len[:, None]) / 6.0
return cand_tok, cand_mask, cand_kill, cand_feat
def make_feat(script: np.ndarray, prot: np.ndarray, spans, t: int) -> np.ndarray:
"""4 通道 token 特徵:script / 在站點 span 內 / 保護段 / 詞界。"""
f = np.zeros((t, 4), np.int64)
f[:, 0] = script
for s, e in spans:
f[min(int(s), t):min(int(e), t), 1] = 1
f[:, 2] = prot
f[1:, 3] = (script[1:] != script[:-1]).astype(np.int64)
return f
|