File size: 4,571 Bytes
e6032a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
"""Phase 4 β€” per-language pyctcdecode decoders backed by KenLM.

Builds a `BeamSearchDecoderCTC` per language and routes by the clip's locale.
Reused by `transcribe.py` (demo) and the Phase 5 evaluation script.
"""
from __future__ import annotations

import os
from typing import Dict, List, Optional


def ctc_labels(tokenizer) -> List[str]:
    """Vocab tokens (id order) mapped to the pyctcdecode/wav2vec2 convention:
    pad β†’ "" (CTC blank), word delimiter "|" β†’ " ", others unchanged.
    """
    vocab = tokenizer.get_vocab()  # {token: id}
    tokens = [tok for tok, _ in sorted(vocab.items(), key=lambda kv: kv[1])]
    labels = []
    for tok in tokens:
        if tok == tokenizer.pad_token:
            labels.append("")          # CTC blank
        elif tok == tokenizer.word_delimiter_token:
            labels.append(" ")         # word boundary
        else:
            labels.append(tok)
    return labels


SPECIAL_TOKENS = ("<unk>", "<s>", "</s>")
MIN_UNIGRAM_COVERAGE = 0.90


def arpa_unigrams(arpa_path: str) -> List[str]:
    """The .arpa's own 1-gram vocabulary (excluding KenLM's special tokens)."""
    words, in_section = [], False
    with open(arpa_path, encoding="utf-8") as f:
        for line in f:
            if not in_section:
                if line.startswith("\\1-grams:"):
                    in_section = True
                continue
            if line.startswith("\\"):     # next section (\2-grams: / \end\)
                break
            parts = line.rstrip("\n").split("\t")   # blank lines -> len 1, skipped
            if len(parts) >= 2 and parts[1] not in SPECIAL_TOKENS:
                words.append(parts[1])
    return words


def check_unigram_coverage(arpa_path: str, unigrams: List[str],
                           min_coverage: float = MIN_UNIGRAM_COVERAGE) -> float:
    """Fraction of the LM's OWN vocabulary present in `unigrams`; raises if low.

    pyctcdecode adds `unk_score_offset` (-10 logp) to every word that is not in
    the unigram list, EVEN IF the KenLM models it well. So a unigram list that
    under-covers the .arpa turns a large LM into a tiny closed vocabulary with a
    -10 cliff around it. pyctcdecode only warns on the OPPOSITE direction
    (unigrams missing from the LM), which stays near 100% precisely when this
    failure is at its worst β€” so it cannot catch a stale list. Checked here.
    """
    lm_vocab = set(arpa_unigrams(arpa_path))
    if not lm_vocab:
        return 1.0
    coverage = len(lm_vocab & set(unigrams)) / len(lm_vocab)
    if coverage < min_coverage:
        raise RuntimeError(
            f"unigrams.txt covers only {coverage:.2%} of {arpa_path}'s "
            f"{len(lm_vocab)} vocabulary words (min {min_coverage:.0%}) β€” it is "
            f"STALE relative to this .arpa. pyctcdecode would penalise the "
            f"other {len(lm_vocab - set(unigrams))} LM-known words by "
            f"unk_score_offset. Regenerate it from the corpus this .arpa was "
            f"trained on:\n"
            f"  python src/lm/build_kenlm.py --unigrams-only <lang> <corpus.txt>"
        )
    return coverage


def load_decoders(
    lm_dir: str,
    languages: List[str],
    labels: List[str],
    alpha: float,
    beta: float,
) -> Dict[str, "object"]:
    """Return {lang: BeamSearchDecoderCTC}. Skips languages with no .arpa."""
    from pyctcdecode import build_ctcdecoder

    decoders: Dict[str, object] = {}
    for lang in languages:
        arpa = os.path.join(lm_dir, lang, f"{lang}.arpa")
        if not os.path.exists(arpa):
            print(f"  [decoder] no LM for '{lang}' ({arpa}); will fall back to greedy")
            continue
        unigrams: Optional[List[str]] = None
        uni = os.path.join(lm_dir, lang, "unigrams.txt")
        if os.path.exists(uni):
            with open(uni, encoding="utf-8") as f:
                unigrams = [w for w in f.read().split("\n") if w]
            cov = check_unigram_coverage(arpa, unigrams)
            print(f"  [decoder] {lang}: unigrams cover {cov:.1%} of LM vocab")
        decoders[lang] = build_ctcdecoder(
            labels, kenlm_model_path=arpa, unigrams=unigrams, alpha=alpha, beta=beta
        )
    return decoders


def decode_logits(decoders, lang, logits, beam_width: int):
    """LM-decode a single (T, vocab) logits array for `lang`. Returns text or None
    if no decoder exists for that language (caller should fall back to greedy)."""
    dec = decoders.get(lang)
    if dec is None:
        return None
    return dec.decode(logits, beam_width=beam_width).strip()