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