| """Shared helpers: dev sentences (walk2 snapshot, read-only), ids + legal mask via sloth's own |
| zhuyin_fmt/slothe_rt (imported read-only; run with PYTHONDONTWRITEBYTECODE=1).""" |
| import json, os, sys |
| import numpy as np |
| W = os.environ.get("MCBPMF_LM_WORK", ".") |
| sys.dont_write_bytecode = True |
| sys.path.insert(0, W + "/sloth") |
| import zhuyin_fmt |
| from slothe_rt import hf, UNK_CHAR_ID |
| SYL = json.load(open(hf("syl_vocab.json"), encoding="utf-8")) |
| MASK = np.load(hf("syl2legal.npz"))["mask"] |
|
|
| def dev_sentences(set_name="dev"): |
| snap = f"{W}/walk2/snap/{set_name}" |
| sents = {json.loads(l)["sent_id"]: json.loads(l) for l in open(snap + "/sentences.jsonl", encoding="utf-8") if l.strip()} |
| walks = [json.loads(l) for l in open(snap + "/walks.jsonl", encoding="utf-8") if l.strip()] |
| return [(w["sent_id"], sents[w["sent_id"]]["readings"]) for w in walks] |
|
|
| def ids_of(readings): |
| _, ids, _ = zhuyin_fmt.readings_to_ids(readings, SYL) |
| return np.asarray(ids, dtype=np.int32) |
|
|
| def logprobs(ids, lg): |
| legal = MASK[ids] |
| z = np.where(legal, lg.astype(np.float32), -np.inf) |
| mx = z.max(1, keepdims=True) |
| return z - (mx + np.log(np.exp(z - mx).sum(1, keepdims=True))), legal |
|
|