"""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", ".") # research workspace root sys.dont_write_bytecode = True sys.path.insert(0, W + "/sloth") import zhuyin_fmt # noqa from slothe_rt import hf, UNK_CHAR_ID # noqa 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) # float32, exactly slothe_rt.SlothE.logprobs mx = z.max(1, keepdims=True) return z - (mx + np.log(np.exp(z - mx).sum(1, keepdims=True))), legal