Download twlat/features.py from JacobLinCool/Twinity-1: direct link, hf CLI and curl.
- Browser
- Download file 7.01 kB
-
https://huggingface.co/JacobLinCool/Twinity-1/resolve/main/twlat/features.py
- Command line
-
hf download hf://JacobLinCool/Twinity-1/twlat/features.py
-
curl -L -o features.py https://huggingface.co/JacobLinCool/Twinity-1/resolve/main/twlat/features.py
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 | |