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()