Recently Written · git

subread-android

SubRead for Android: times an audiobook against its ebook on the device

git clone https://github.com/equwal/subread-android

Log | Files | Refs


tools/make_golden.py (8521 bytes)

1 """Generate golden fixtures for the Kotlin aligner from the reference Python one.
2 
3 The Android app ports subplz's alignment (ats.align + subplz.align.shift_align)
4 to Kotlin. This runs the *original* on real Whisper-tiny transcripts and records
5 every intermediate stage, so each ported function can be checked against the
6 real thing with identical inputs - not just the end result.
7 
8 Needs the subplz environment (ats, Bio, rapidfuzz):
9 
10     SubPlz/venv311/Scripts/python tools/make_golden.py --cache <dir> --book <name> ...
11 
12 Fixtures committed to the repo use public-domain text only (Aozora Bunko).
13 Anything else goes to golden-local/, which is gitignored; tests skip it when
14 it is absent.
15 """
16 
17 from __future__ import annotations
18 
19 import argparse
20 import ast
21 import json
22 import re
23 import sys
24 import zipfile
25 from pathlib import Path
26 
27 from ats import align as ats_align
28 from ats.lang import get_lang
29 from ats.main import to_subs
30 from Bio import Align
31 from subplz.align import shift_align
32 from subplz.cli import END_PUNC, START_PUNC
33 
34 OTHER_PUNC = "* ,,、…"
35 PREPEND = START_PUNC
36 APPEND = END_PUNC + OTHER_PUNC
37 NOPEND = (
38     "うぁぃぅぇぉっゃゅょゎゕゖァィゥェォヵㇰヶㇱㇲッㇳㇴㇵㇶㇷㇷ゚ㇸㇹㇺャュョㇻㇼㇽㇾㇿヮ… \x20"
39 )
40 
41 
42 def plain(x):
43     """Deep copy as plain JSON types. align_sub leaves numpy ints in its lists."""
44     return json.loads(json.dumps(x, default=int))
45 
46 
47 class Para:
48     """Duck-types ats.main.Paragraph: to_subs only ever calls .text()."""
49 
50     def __init__(self, s: str):
51         self._s = s
52 
53     def text(self) -> str:
54         return self._s
55 
56 
57 def load_transcript(path: Path) -> dict:
58     return ast.literal_eval(path.read_text(encoding="utf-8"))
59 
60 
61 def aozora_paragraphs(zip_path: Path) -> list[str]:
62     """Body paragraphs of an Aozora Bunko ruby text, markup removed."""
63     with zipfile.ZipFile(zip_path) as z:
64         name = next(n for n in z.namelist() if n.lower().endswith(".txt"))
65         raw = z.read(name).decode("shift_jis", errors="replace")
66     lines = raw.replace("\r\n", "\n").split("\n")
67 
68     # Header: title block, then a legend fenced by two dashed rules.
69     rules = [i for i, l in enumerate(lines) if l.startswith("-----")]
70     body = lines[rules[1] + 1 :] if len(rules) >= 2 else lines
71     # Footer: colophon.
72     for i, l in enumerate(body):
73         if l.startswith("底本:"):
74             body = body[:i]
75             break
76 
77     out = []
78     for l in body:
79         l = re.sub(r"[#[^]]*]", "", l)      # editorial annotations
80         l = re.sub(r"《[^》]*》", "", l)        # ruby readings
81         l = l.replace("|", "").strip()         # ruby base marker
82         if l:
83             out.append(l)
84     return out
85 
86 
87 def clean_len(lang, s: str) -> int:
88     return len(lang.clean(s))
89 
90 
91 def run_case(name: str, language: str, segments: list[dict], paragraphs: list[str]) -> dict:
92     lang = get_lang(language)  # same call do_batch makes: no punctuation args
93     transcript = [s["text"] for s in segments]
94 
95     # Re-run the aligner by hand as well as through ats.align.align, to capture
96     # the raw coordinates and the pre-heuristic segments.
97     aligner = Align.PairwiseAligner(
98         mode="global", match_score=1, open_gap_score=-0.8,
99         mismatch_score=-0.6, extend_gap_score=-0.5,
100     )
101     t_clean = [lang.clean(i) for i in transcript]
102     p_clean = [lang.clean(i) for i in paragraphs]
103     best = aligner.align("".join(p_clean), "".join(t_clean))[0]
104     coords = best.coordinates
105     # align_sub writes into this array (it keeps a numpy view of a column and
106     # assigns through it), so take the copy that gets recorded first.
107     coords_in = [[int(v) for v in row] for row in coords]
108 
109     raw_segments = ats_align.align_sub(coords, p_clean, t_clean)
110     after_sub = plain(raw_segments)
111     ats_align.fix(lang, paragraphs, p_clean, raw_segments)
112     after_fix = plain(raw_segments)
113     ats_align.fix_punc(paragraphs, raw_segments, set(PREPEND), set(APPEND), set(NOPEND))
114     after_punc = plain(raw_segments)
115 
116     # The public entry point must agree with the hand-run stages.
117     alignment, _ = ats_align.align(
118         None, lang, transcript, paragraphs, [], set(PREPEND), set(APPEND), set(NOPEND)
119     )
120     assert plain(alignment) == after_punc, "stage capture drifted"
121 
122     subs = to_subs([Para(p) for p in paragraphs], segments, alignment, 0, None)
123     cues_raw = [{"text": s.text, "start": s.start, "end": s.end} for s in subs]
124     shifted = shift_align(subs)
125     cues = [{"text": s.text, "start": s.start, "end": s.end} for s in shifted]
126 
127     matched = sum(1 for c in cues_raw if not c["text"].startswith("*"))
128     print(f"{name}: {len(segments)} segs, {len(paragraphs)} paras, "
129           f"{sum(map(len, t_clean))}x{sum(map(len, p_clean))} chars, "
130           f"score {best.score:.1f}, {matched}/{len(cues_raw)} cues matched")
131 
132     return {
133         "name": name,
134         "language": language,
135         "prepend": PREPEND, "append": APPEND, "nopend": NOPEND,
136         "transcript": [{"text": s["text"], "start": s["start"], "end": s["end"]}
137                        for s in segments],
138         "paragraphs": paragraphs,
139         "transcript_clean": t_clean,
140         "paragraphs_clean": p_clean,
141         "score": float(best.score),
142         "coords": coords_in,
143         "after_align_sub": after_sub,
144         "after_fix": after_fix,
145         "after_fix_punc": after_punc,
146         "cues_raw": cues_raw,
147         "cues": cues,
148     }
149 
150 
151 def chapter_files(cache: Path, book: str) -> list[Path]:
152     rx = re.compile(re.escape(book) + r"\.(\d+)\.tiny\.subs$")
153     found = [(int(m.group(1)), p) for p in cache.iterdir() if (m := rx.search(p.name))]
154     return [p for _, p in sorted(found)]
155 
156 
157 def main() -> None:
158     sys.stdout.reconfigure(encoding="utf-8")
159     ap = argparse.ArgumentParser()
160     ap.add_argument("--cache", type=Path, required=True, help="subplz transcript cache dir")
161     ap.add_argument("--book", required=True, help="audio file name the cache entries start with")
162     ap.add_argument("--aozora", type=Path, help="Aozora Bunko ruby zip with the matching text")
163     ap.add_argument("--epub", type=Path, help="epub with the matching text")
164     ap.add_argument("--out", type=Path, required=True)
165     ap.add_argument("--prefix", required=True)
166     ap.add_argument("--spans", default="0:3,3:6,20:24",
167                     help="audio chapter ranges to turn into cases, e.g. 0:3,10:14")
168     args = ap.parse_args()
169 
170     files = chapter_files(args.cache, args.book)
171     if not files:
172         sys.exit(f"no cache entries for {args.book!r} in {args.cache}")
173     transcripts = [load_transcript(p) for p in files]
174     language = transcripts[0]["language"]
175     lang = get_lang(language)
176 
177     if args.aozora:
178         paragraphs = aozora_paragraphs(args.aozora)
179     else:
180         from ats.main import Epub
181         paragraphs = [p.text() for ch in Epub.from_file(str(args.epub)) for p in ch.text()]
182         paragraphs = [p for p in paragraphs if p.strip()]
183 
184     # Where each audio chapter starts in the text, by cumulative cleaned length.
185     # Approximate on purpose: a fixture whose text runs a little long or short
186     # at either end is a more honest test than a perfectly trimmed one.
187     para_cum = [0]
188     for p in paragraphs:
189         para_cum.append(para_cum[-1] + clean_len(lang, p))
190     chap_cum = [0]
191     for t in transcripts:
192         chap_cum.append(chap_cum[-1] + sum(clean_len(lang, s["text"]) for s in t["segments"]))
193     scale = para_cum[-1] / max(1, chap_cum[-1])
194 
195     def para_at(chars: float) -> int:
196         target = chars * scale
197         return next((i for i, c in enumerate(para_cum) if c >= target), len(paragraphs))
198 
199     args.out.mkdir(parents=True, exist_ok=True)
200     for span in args.spans.split(","):
201         a, b = (int(x) for x in span.split(":"))
202         b = min(b, len(transcripts))
203         segments, offset = [], 0.0
204         for t in transcripts[a:b]:
205             for s in t["segments"]:
206                 segments.append({"text": s["text"], "start": s["start"] + offset,
207                                  "end": s["end"] + offset})
208             offset = segments[-1]["end"] if segments else offset
209         lo = max(0, para_at(chap_cum[a]) - 2)
210         hi = min(len(paragraphs), para_at(chap_cum[b]) + 2)
211         name = f"{args.prefix}_{a:03d}_{b:03d}"
212         case = run_case(name, language, segments, paragraphs[lo:hi])
213         (args.out / f"{name}.json").write_text(
214             json.dumps(case, ensure_ascii=False, indent=0), encoding="utf-8")
215 
216 
217 if __name__ == "__main__":
218     main()