File size: 1,282 Bytes
ff5f59d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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