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