Download twlat/runtime_v3.py from JacobLinCool/Twinity-1: direct link, hf CLI and curl.
- Browser
- Download file 10.4 kB
-
https://huggingface.co/JacobLinCool/Twinity-1/resolve/main/twlat/runtime_v3.py
- Command line
-
hf download hf://JacobLinCool/Twinity-1/twlat/runtime_v3.py
-
curl -L -o runtime_v3.py https://huggingface.co/JacobLinCool/Twinity-1/resolve/main/twlat/runtime_v3.py
10.4 kB
| """TWLAT V3 推論管線。 | |
| 輸入 → protect(PUA 遮罩保護段)→ safe_normalize(twv) | |
| → lattice(字典硬過濾在此發生)→ 無邊 → fast path | |
| → clean/masked 雙 pass forward → Viterbi + τ → 最小 splice | |
| → restore → 輸出 | |
| 超長輸入:滑動視窗(stride 384),每條邊由「它最居中」的視窗評分—— | |
| 不再有 V2 的 INPUT_TOO_LONG 降級路徑。 | |
| """ | |
| from __future__ import annotations | |
| import pathlib | |
| from dataclasses import dataclass, field | |
| import numpy as np | |
| import torch | |
| from twlat import decoder # noqa: E402 | |
| from twlat.features import (FEAT_DIM, L_MAX, LexTables, # noqa: E402 | |
| assemble_cands, make_feat, site_arrays, | |
| text_arrays) | |
| from twlat.lattice import LatticeBuilder, build_masked_view # noqa: E402 | |
| from twlat.model_v3 import TWLATV3, make_config # noqa: E402 | |
| from twlat.normalize import safe_normalize # noqa: E402 | |
| from twlat.paths import data_file # noqa: E402 | |
| from twlat.protect import protect, restore # noqa: E402 | |
| import json # noqa: E402 | |
| WIN, STRIDE = 512, 384 | |
| class ResultV3: | |
| output: str | |
| decisions: list = field(default_factory=list) | |
| fast_path: bool = False | |
| error: str | None = None | |
| class TWLATV3Runtime: | |
| def __init__(self, ckpt: str, device: str | None = None, | |
| tau: dict[str, float] | None = None, | |
| lexicon_path=None, fo_bonus: float = 0.0): | |
| lex_p = lexicon_path or data_file("dict/lattice_lexicon.json") | |
| self.lb = LatticeBuilder(lex_p) | |
| self.lex = LexTables(lex_p) | |
| self.vocab = json.loads( | |
| data_file("dict/char_vocab_v3.json").read_text(encoding="utf-8")) | |
| self.fo_bonus = fo_bonus | |
| tp = data_file("model/tau.json") | |
| self.tau = tau if tau is not None else ( | |
| json.loads(tp.read_text()) if tp.exists() else decoder.DEFAULT_TAU) | |
| self.device = torch.device( | |
| device or ("cuda" if torch.cuda.is_available() | |
| else "mps" if torch.backends.mps.is_available() else "cpu")) | |
| ck = torch.load(ckpt, map_location="cpu", weights_only=False) | |
| self.model = TWLATV3(make_config(**ck["mcfg"])) | |
| self.model.load_state_dict(ck["model"]) | |
| self.model.eval().to(self.device) | |
| # observed 相依特徵(is_keep/長度差)只有 finetune 階段見過; | |
| # 對 pretrain-only checkpoint 餵入會是分佈外輸入 | |
| self.reveal = ck.get("phase") == "finetune" | |
| ck_lex = ck.get("lexicon_version") | |
| if ck_lex and ck_lex != self.lb.version: | |
| # 熱更新是功能不是錯誤:記錄但不擋(held-out 評測依賴此路徑) | |
| self.lexicon_mismatch = (ck_lex, self.lb.version) | |
| else: | |
| self.lexicon_mismatch = None | |
| # ------------------------------------------------------------------ # | |
| def score_batch(self, texts: list[str], batch_size: int = 32): | |
| """評分階段:→ list[(base, sv, lat, logits[n_edges,C] | None)]。 | |
| τ 校準重用此結果做多組解碼,不必重跑模型。""" | |
| stages = [] | |
| jobs = [] | |
| for ti, t in enumerate(texts): | |
| q, sv = protect(t) | |
| base = safe_normalize(q, do_twv=True) | |
| lat = self.lb.build(base) | |
| stages.append([base, sv, lat, None]) | |
| if not lat.edges: | |
| continue | |
| wins = self._windows(len(base)) | |
| assign = self._assign(lat, wins) | |
| for wi, (ws, we) in enumerate(wins): | |
| eixs = assign[wi] | |
| if not eixs: | |
| continue | |
| jobs.append(self._make_job(ti, base[ws:we], ws, | |
| [lat.edges[i] for i in eixs], eixs)) | |
| logit_map: dict[int, dict[int, np.ndarray]] = {} | |
| jobs.sort(key=lambda j: len(j["text"])) | |
| for s in range(0, len(jobs), batch_size): | |
| chunk = jobs[s:s + batch_size] | |
| out = self.model(self._stack(chunk)) | |
| lg = out["cand_logits"].float().cpu().numpy() | |
| for bi, j in enumerate(chunk): | |
| m = logit_map.setdefault(j["ti"], {}) | |
| for k, eix in enumerate(j["eixs"]): | |
| m[eix] = lg[bi, k] | |
| for ti, st in enumerate(stages): | |
| lat = st[2] | |
| if not lat.edges: | |
| continue | |
| lm = logit_map.get(ti, {}) | |
| cmax = max((v.shape[0] for v in lm.values()), default=1) | |
| logits = np.full((len(lat.edges), cmax), -np.inf, np.float32) | |
| for eix, row in lm.items(): | |
| logits[eix, :len(row)] = row | |
| for i, e in enumerate(lat.edges): | |
| if i not in lm: # 不應發生:未覆蓋 → keep | |
| logits[i, e.obs_ix] = 0.0 | |
| st[3] = logits | |
| return stages | |
| def decode_stage(self, stage, tau=None) -> ResultV3: | |
| base, sv, lat, logits = stage | |
| if logits is None: | |
| return ResultV3(output=restore(base, sv), fast_path=True) | |
| edits = decoder.decode(self.lb, lat, logits, | |
| self.tau if tau is None else tau, | |
| fo_bonus=self.fo_bonus) | |
| return ResultV3( | |
| output=restore(decoder.splice(base, edits), sv), | |
| decisions=[{"span": [e.start, e.end], "from": e.observed, | |
| "to": e.replacement, "utility": round(e.utility, 4), | |
| "rule_type": e.rule_type} for e in edits]) | |
| def convert_batch(self, texts: list[str], batch_size: int = 32 | |
| ) -> list[ResultV3]: | |
| return [self.decode_stage(s) | |
| for s in self.score_batch(texts, batch_size)] | |
| # ------------------------------------------------------------------ # | |
| def _windows(n: int) -> list[tuple[int, int]]: | |
| if n <= WIN: | |
| return [(0, n)] | |
| wins, s = [], 0 | |
| while True: | |
| wins.append((s, min(s + WIN, n))) | |
| if s + WIN >= n: | |
| return wins | |
| s += STRIDE | |
| def _assign(lat, wins) -> dict[int, list[int]]: | |
| """每條邊 → 它最居中的視窗。""" | |
| out: dict[int, list[int]] = {wi: [] for wi in range(len(wins))} | |
| for i, e in enumerate(lat.edges): | |
| best_wi, best_d = None, None | |
| for wi, (ws, we) in enumerate(wins): | |
| if e.start >= ws and e.end <= we: | |
| c = (ws + we) / 2 | |
| d = abs((e.start + e.end) / 2 - c) | |
| if best_d is None or d < best_d: | |
| best_wi, best_d = wi, d | |
| if best_wi is not None: | |
| out[best_wi].append(i) | |
| return out | |
| def _make_job(self, ti, text, offset, edges, eixs): | |
| ids, script, prot = text_arrays(text, self.vocab) | |
| # 邊座標平移到視窗座標 | |
| shifted = [] | |
| for e in edges: | |
| se = type(e)(e.start - offset, e.end - offset, e.gid, e.obs_ix, | |
| e.cand_kill, e.word_crossing, e.word_contained) | |
| shifted.append(se) | |
| sa = site_arrays(self.lb, shifted, text) | |
| mids, m_span, _ = build_masked_view(ids, sa["span"], sa["maskable"]) | |
| mscript, _, _ = build_masked_view(script, sa["span"], sa["maskable"], | |
| mask_id=4) | |
| mprot, _, _ = build_masked_view(prot.astype(np.uint8), sa["span"], | |
| sa["maskable"], mask_id=0) | |
| cand_tok, cand_mask, cand_kill, cand_feat = assemble_cands( | |
| self.lex, sa["gid"], sa["obs"], sa["clue"], sa["eng"], | |
| sa["flags"], sa["kill"], reveal_observed=self.reveal) | |
| return {"ti": ti, "eixs": eixs, "text": text, | |
| "ids": ids, "script": script, "prot": prot, | |
| "feat": make_feat(script, prot, sa["span"], len(text)), | |
| "mids": mids, "m_span": m_span.astype(np.int64), | |
| "mfeat": make_feat(mscript, mprot, m_span, len(mids)), | |
| "c_span": sa["span"], "obs": sa["obs"], | |
| "cand_tok": cand_tok, "cand_mask": cand_mask, | |
| "cand_kill": cand_kill, "cand_feat": cand_feat} | |
| def _stack(self, jobs): | |
| B = len(jobs) | |
| T = max(len(j["ids"]) for j in jobs) | |
| Tm = max(len(j["mids"]) for j in jobs) | |
| S = max(len(j["c_span"]) for j in jobs) | |
| C = max(j["cand_tok"].shape[1] for j in jobs) | |
| out = { | |
| "ids": np.zeros((B, T), np.int64), | |
| "feat": np.zeros((B, T, 4), np.int64), | |
| "pad": np.zeros((B, T), bool), | |
| "mids": np.zeros((B, Tm), np.int64), | |
| "mfeat": np.zeros((B, Tm, 4), np.int64), | |
| "mpad": np.zeros((B, Tm), bool), | |
| "c_span": np.zeros((B, S, 2), np.int64), | |
| "m_span": np.zeros((B, S, 2), np.int64), | |
| "site_mask": np.zeros((B, S), bool), | |
| "cand_tok": np.zeros((B, S, C, L_MAX), np.int64), | |
| "cand_mask": np.zeros((B, S, C), bool), | |
| "cand_kill": np.zeros((B, S, C), bool), | |
| "cand_feat": np.zeros((B, S, C, FEAT_DIM), np.float32), | |
| } | |
| for b, j in enumerate(jobs): | |
| n, tm = len(j["ids"]), len(j["mids"]) | |
| ns, c = len(j["c_span"]), j["cand_tok"].shape[1] | |
| out["ids"][b, :n] = j["ids"] | |
| out["feat"][b, :n] = j["feat"] | |
| out["pad"][b, :n] = True | |
| out["mids"][b, :tm] = j["mids"] | |
| out["mfeat"][b, :tm] = j["mfeat"] | |
| out["mpad"][b, :tm] = True | |
| out["c_span"][b, :ns] = j["c_span"] | |
| out["m_span"][b, :ns] = j["m_span"] | |
| out["site_mask"][b, :ns] = True | |
| out["cand_tok"][b, :ns, :c] = j["cand_tok"] | |
| out["cand_mask"][b, :ns, :c] = j["cand_mask"] | |
| out["cand_kill"][b, :ns, :c] = j["cand_kill"] | |
| out["cand_feat"][b, :ns, :c] = j["cand_feat"] | |
| return {k: torch.from_numpy(v).to(self.device) for k, v in out.items()} | |