altasr 1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- altasr/__init__.py +92 -0
- altasr/audio.py +208 -0
- altasr/benchmark.py +374 -0
- altasr/config.py +245 -0
- altasr/export.py +152 -0
- altasr/inference/__init__.py +44 -0
- altasr/inference/checkpoint.py +111 -0
- altasr/inference/decoding.py +304 -0
- altasr/inference/evaluate.py +194 -0
- altasr/inference/pool.py +159 -0
- altasr/inference/streaming.py +247 -0
- altasr/inference/transcribe.py +88 -0
- altasr/inference/transcriber.py +331 -0
- altasr/integrations/__init__.py +3 -0
- altasr/integrations/llm.py +231 -0
- altasr/lm.py +266 -0
- altasr/metrics.py +126 -0
- altasr/model/__init__.py +29 -0
- altasr/model/attention.py +104 -0
- altasr/model/block.py +66 -0
- altasr/model/convolution.py +59 -0
- altasr/model/ctc.py +104 -0
- altasr/model/encoder.py +109 -0
- altasr/model/feedforward.py +27 -0
- altasr/model/positional.py +75 -0
- altasr/model/subsampling.py +40 -0
- altasr/text.py +313 -0
- altasr/training/__init__.py +36 -0
- altasr/training/data.py +329 -0
- altasr/training/finetune.py +226 -0
- altasr/training/prepare.py +101 -0
- altasr/training/train.py +248 -0
- altasr/training/trainer.py +488 -0
- altasr/utils.py +154 -0
- altasr-1.dist-info/METADATA +428 -0
- altasr-1.dist-info/RECORD +39 -0
- altasr-1.dist-info/WHEEL +5 -0
- altasr-1.dist-info/entry_points.txt +9 -0
- altasr-1.dist-info/top_level.txt +1 -0
altasr/__init__.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
"""
|
|
2
|
+
ALTASR -- Adaptive Low-latency Transcription ASR for Kinyarwanda.
|
|
3
|
+
==================================================================
|
|
4
|
+
|
|
5
|
+
A clear, extensible Conformer-CTC speech-recognition package, organised as
|
|
6
|
+
three subpackages under one roof:
|
|
7
|
+
|
|
8
|
+
``altasr.model`` the network -- one file per architectural component.
|
|
9
|
+
``altasr.training`` everything to train: dataset/preprocessing, the
|
|
10
|
+
epoch-based trainer, feature precomputation, the
|
|
11
|
+
``altasr-train`` CLI.
|
|
12
|
+
``altasr.inference`` everything to use a trained model: checkpoint loading,
|
|
13
|
+
the :class:`Transcriber` API, ``altasr-transcribe`` /
|
|
14
|
+
``altasr-evaluate`` CLIs.
|
|
15
|
+
|
|
16
|
+
Shared foundations live at the package root: ``config`` (one dataclass tree,
|
|
17
|
+
JSON in/out), ``text`` (normalisation + tokenizer), ``audio`` (loading,
|
|
18
|
+
log-mels, SpecAugment) and ``metrics`` (WER/CER).
|
|
19
|
+
|
|
20
|
+
Quick use of a published checkpoint::
|
|
21
|
+
|
|
22
|
+
from altasr import Transcriber
|
|
23
|
+
asr = Transcriber.from_pretrained("path/to/checkpoint") # folder or .pt
|
|
24
|
+
print(asr.transcribe("recording.mp3"))
|
|
25
|
+
|
|
26
|
+
Quick training (after ``pip install altasr`` or ``pip install -e .``)::
|
|
27
|
+
|
|
28
|
+
altasr-train --preset medium --audio-root /data/track_b \
|
|
29
|
+
--train /data/track_b/train.json --val /data/track_b/dev_test.json
|
|
30
|
+
|
|
31
|
+
Heavy dependencies (torch) are imported lazily: ``import altasr`` plus the
|
|
32
|
+
config/text/metrics layers work on any machine; the model, training and
|
|
33
|
+
inference layers require PyTorch.
|
|
34
|
+
"""
|
|
35
|
+
from __future__ import annotations
|
|
36
|
+
|
|
37
|
+
__version__ = "0.5.1"
|
|
38
|
+
|
|
39
|
+
# Torch-free foundations -- always importable.
|
|
40
|
+
from .config import Config, PRESETS, load_config, preset # noqa: F401
|
|
41
|
+
from .text import (CharTokenizer, BPETokenizer, normalize_text, # noqa: F401
|
|
42
|
+
strip_punctuation, build_tokenizer, load_tokenizer) # noqa: F401
|
|
43
|
+
from .metrics import cer, corpus_wer_cer, corpus_metrics, wer # noqa: F401
|
|
44
|
+
|
|
45
|
+
# Torch-backed public API, resolved lazily (PEP 562) so `import altasr`
|
|
46
|
+
# works without PyTorch and errors stay clear when torch is missing.
|
|
47
|
+
_LAZY = {
|
|
48
|
+
"Transcriber": "altasr.inference.transcriber",
|
|
49
|
+
"load_checkpoint": "altasr.inference.checkpoint",
|
|
50
|
+
"save_checkpoint": "altasr.inference.checkpoint",
|
|
51
|
+
"build_model": "altasr.model.ctc",
|
|
52
|
+
"ConformerCTC": "altasr.model.ctc",
|
|
53
|
+
"Trainer": "altasr.training.trainer",
|
|
54
|
+
"ASRDataset": "altasr.training.data",
|
|
55
|
+
"load_records": "altasr.training.data",
|
|
56
|
+
"StreamingSession": "altasr.inference.streaming",
|
|
57
|
+
"TranscriberPool": "altasr.inference.pool",
|
|
58
|
+
"resolve_devices": "altasr.inference.pool",
|
|
59
|
+
"AudioDecodeError": "altasr.audio",
|
|
60
|
+
"Hotwords": "altasr.inference.decoding",
|
|
61
|
+
"TranscriptionResult": "altasr.inference.decoding",
|
|
62
|
+
"ctc_prefix_beam_search": "altasr.inference.decoding",
|
|
63
|
+
"NGramLM": "altasr.lm",
|
|
64
|
+
"Lexicon": "altasr.lm",
|
|
65
|
+
"build_decoding_resources": "altasr.lm",
|
|
66
|
+
"load_decoding_resources": "altasr.lm",
|
|
67
|
+
"LLMCorrector": "altasr.integrations.llm",
|
|
68
|
+
"SmartTranscriber": "altasr.integrations.llm",
|
|
69
|
+
"ASR_TOOL_SCHEMA": "altasr.integrations.llm",
|
|
70
|
+
"handle_tool_call": "altasr.integrations.llm",
|
|
71
|
+
"transcribe_stream": "altasr.inference.streaming",
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def __getattr__(name):
|
|
76
|
+
target = _LAZY.get(name)
|
|
77
|
+
if target is None:
|
|
78
|
+
raise AttributeError(f"module 'altasr' has no attribute {name!r}")
|
|
79
|
+
import importlib
|
|
80
|
+
try:
|
|
81
|
+
module = importlib.import_module(target)
|
|
82
|
+
except ImportError as exc:
|
|
83
|
+
raise ImportError(
|
|
84
|
+
f"altasr.{name} requires PyTorch. Install the backend with "
|
|
85
|
+
f"`pip install torch torchaudio` ({exc})") from exc
|
|
86
|
+
value = getattr(module, name)
|
|
87
|
+
globals()[name] = value
|
|
88
|
+
return value
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def __dir__():
|
|
92
|
+
return sorted(list(globals()) + list(_LAZY))
|
altasr/audio.py
ADDED
|
@@ -0,0 +1,208 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Audio loading and feature extraction.
|
|
3
|
+
|
|
4
|
+
Loading
|
|
5
|
+
-------
|
|
6
|
+
``load_audio`` returns a mono float32 waveform at the target sample rate. It
|
|
7
|
+
tries decoders in a robust order with no system-FFmpeg dependency:
|
|
8
|
+
|
|
9
|
+
1. **soundfile** (libsndfile): WAV / FLAC / OGG -- fast and reliable;
|
|
10
|
+
2. **PyAV** (FFmpeg bundled inside the pip wheel): MP3 / Opus / M4A / WebM and
|
|
11
|
+
extension-less web recordings;
|
|
12
|
+
3. **torchaudio.load** as a last resort.
|
|
13
|
+
|
|
14
|
+
Resampling uses ``torchaudio.functional.resample`` (windowed-sinc).
|
|
15
|
+
|
|
16
|
+
Features
|
|
17
|
+
--------
|
|
18
|
+
``LogMel`` converts a waveform to a log mel spectrogram:
|
|
19
|
+
|
|
20
|
+
STFT (n_fft=400 = 25 ms window, hop=160 = 10 ms) -> power spectrum
|
|
21
|
+
-> mel filterbank (n_mels=80) -> log(mel + 1e-6)
|
|
22
|
+
-> per-utterance normalisation: (x - mean) / (std + 1e-5)
|
|
23
|
+
|
|
24
|
+
Per-utterance mean/variance normalisation makes the network robust to channel
|
|
25
|
+
loudness/recording differences and lets SpecAugment mask with zeros (zero ==
|
|
26
|
+
the mean after normalisation).
|
|
27
|
+
|
|
28
|
+
``SpecAugment`` (Park et al., 2019) randomly zeroes frequency bands and time
|
|
29
|
+
spans **during training only** -- the single most effective ASR regulariser.
|
|
30
|
+
The implementation is fully vectorised over the batch (no Python loops).
|
|
31
|
+
"""
|
|
32
|
+
from __future__ import annotations
|
|
33
|
+
|
|
34
|
+
from typing import Optional
|
|
35
|
+
|
|
36
|
+
import torch
|
|
37
|
+
import torch.nn as nn
|
|
38
|
+
|
|
39
|
+
from .config import AudioConfig
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
# --------------------------------------------------------------------------- #
|
|
43
|
+
# Loading
|
|
44
|
+
# --------------------------------------------------------------------------- #
|
|
45
|
+
class AudioDecodeError(RuntimeError):
|
|
46
|
+
"""Raised when every decoder failed for a file. Carries the per-decoder
|
|
47
|
+
reasons so the dataset layer can log a useful one-line warning and skip
|
|
48
|
+
the file instead of crashing an hours-long training run."""
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def load_audio(path: str, target_sr: int = 16_000) -> torch.Tensor:
|
|
52
|
+
"""Decode ``path`` to a mono float32 waveform (T,) at ``target_sr`` Hz.
|
|
53
|
+
|
|
54
|
+
Decoders are tried in order (soundfile -> PyAV -> torchaudio); each
|
|
55
|
+
failure is recorded, and only if ALL fail is :class:`AudioDecodeError`
|
|
56
|
+
raised. The torchaudio fallback is fully wrapped: on torchaudio >= 2.9
|
|
57
|
+
it delegates to torchcodec, which hard-crashes with a RuntimeError when
|
|
58
|
+
system FFmpeg libraries are missing -- that must never take down a
|
|
59
|
+
training job when soundfile/PyAV simply didn't support the format.
|
|
60
|
+
"""
|
|
61
|
+
wav, sr = None, None
|
|
62
|
+
errors = []
|
|
63
|
+
|
|
64
|
+
try: # 1) soundfile
|
|
65
|
+
import soundfile as sf
|
|
66
|
+
data, sr = sf.read(path, dtype="float32", always_2d=True)
|
|
67
|
+
wav = torch.from_numpy(data.mean(axis=1))
|
|
68
|
+
except Exception as exc:
|
|
69
|
+
errors.append(f"soundfile: {exc}")
|
|
70
|
+
|
|
71
|
+
if wav is None: # 2) PyAV
|
|
72
|
+
try:
|
|
73
|
+
import av
|
|
74
|
+
import numpy as np
|
|
75
|
+
chunks = []
|
|
76
|
+
with av.open(path) as container:
|
|
77
|
+
stream = container.streams.audio[0]
|
|
78
|
+
resampler = av.audio.resampler.AudioResampler(
|
|
79
|
+
format="fltp", layout="mono", rate=target_sr)
|
|
80
|
+
for frame in container.decode(stream):
|
|
81
|
+
for rs in resampler.resample(frame):
|
|
82
|
+
chunks.append(rs.to_ndarray().reshape(-1))
|
|
83
|
+
if not chunks:
|
|
84
|
+
raise ValueError("no audio frames decoded")
|
|
85
|
+
wav = torch.from_numpy(np.concatenate(chunks).astype("float32"))
|
|
86
|
+
sr = target_sr
|
|
87
|
+
except Exception as exc:
|
|
88
|
+
errors.append(f"pyav: {exc}")
|
|
89
|
+
|
|
90
|
+
if wav is None: # 3) torchaudio
|
|
91
|
+
try:
|
|
92
|
+
import torchaudio
|
|
93
|
+
data, sr = torchaudio.load(path)
|
|
94
|
+
wav = data.mean(dim=0)
|
|
95
|
+
except Exception as exc:
|
|
96
|
+
# torchaudio 2.9+ routes through torchcodec, which raises a huge
|
|
97
|
+
# RuntimeError when FFmpeg shared libraries are absent; compress
|
|
98
|
+
# it to one line so warnings stay readable.
|
|
99
|
+
msg = str(exc).splitlines()[0]
|
|
100
|
+
errors.append(f"torchaudio: {msg}")
|
|
101
|
+
|
|
102
|
+
if wav is None or wav.numel() == 0:
|
|
103
|
+
raise AudioDecodeError(
|
|
104
|
+
f"could not decode {path!r} with any backend: "
|
|
105
|
+
+ " | ".join(errors or ["file empty"]))
|
|
106
|
+
|
|
107
|
+
if sr != target_sr:
|
|
108
|
+
import torchaudio.functional as AF
|
|
109
|
+
wav = AF.resample(wav, orig_freq=int(sr), new_freq=target_sr)
|
|
110
|
+
return wav.contiguous()
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
# --------------------------------------------------------------------------- #
|
|
114
|
+
# Log-mel front-end
|
|
115
|
+
# --------------------------------------------------------------------------- #
|
|
116
|
+
class LogMel(nn.Module):
|
|
117
|
+
"""Waveform (B, S) or (S,) -> normalised log-mel features (B, T, n_mels)."""
|
|
118
|
+
|
|
119
|
+
def __init__(self, cfg: AudioConfig) -> None:
|
|
120
|
+
super().__init__()
|
|
121
|
+
import torchaudio
|
|
122
|
+
self.cfg = cfg
|
|
123
|
+
self.mel = torchaudio.transforms.MelSpectrogram(
|
|
124
|
+
sample_rate=cfg.sample_rate, n_fft=cfg.n_fft,
|
|
125
|
+
hop_length=cfg.hop_length, n_mels=cfg.n_mels, power=2.0,
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
def forward(self, wav: torch.Tensor) -> torch.Tensor:
|
|
129
|
+
if wav.dim() == 1:
|
|
130
|
+
wav = wav.unsqueeze(0)
|
|
131
|
+
mel = self.mel(wav) # (B, n_mels, T)
|
|
132
|
+
logmel = torch.log(mel + 1e-6).transpose(1, 2) # (B, T, n_mels)
|
|
133
|
+
mean = logmel.mean(dim=(1, 2), keepdim=True)
|
|
134
|
+
std = logmel.std(dim=(1, 2), keepdim=True)
|
|
135
|
+
return (logmel - mean) / (std + 1e-5)
|
|
136
|
+
|
|
137
|
+
def num_frames(self, num_samples: torch.Tensor) -> torch.Tensor:
|
|
138
|
+
"""Mel frame count for a sample count: floor(S / hop) + 1 (center pad)."""
|
|
139
|
+
return num_samples // self.cfg.hop_length + 1
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def normalize_logmel(logmel: torch.Tensor) -> torch.Tensor:
|
|
143
|
+
"""Per-utterance (mean, std) normalisation for precomputed features (T, F)."""
|
|
144
|
+
return (logmel - logmel.mean()) / (logmel.std() + 1e-5)
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
# --------------------------------------------------------------------------- #
|
|
148
|
+
# SpecAugment
|
|
149
|
+
# --------------------------------------------------------------------------- #
|
|
150
|
+
class SpecAugment(nn.Module):
|
|
151
|
+
"""Vectorised frequency/time masking of (B, T, F) features; train-time only."""
|
|
152
|
+
|
|
153
|
+
def __init__(self, cfg: AudioConfig) -> None:
|
|
154
|
+
super().__init__()
|
|
155
|
+
self.cfg = cfg
|
|
156
|
+
|
|
157
|
+
@staticmethod
|
|
158
|
+
def _bands(n: int, max_w: torch.Tensor, length: int, b: int,
|
|
159
|
+
device: torch.device) -> torch.Tensor:
|
|
160
|
+
"""(B, length) bool mask, True inside ``n`` random bands per row;
|
|
161
|
+
per-row band width ~ U{0..max_w}, start ~ U{0..length-width}."""
|
|
162
|
+
max_w = max_w.clamp(0, length - 1).view(b, 1).float()
|
|
163
|
+
widths = (torch.rand(b, n, device=device) * (max_w + 1)).long()
|
|
164
|
+
starts = (torch.rand(b, n, device=device)
|
|
165
|
+
* (length - widths + 1).float()).long()
|
|
166
|
+
ar = torch.arange(length, device=device).view(1, 1, length)
|
|
167
|
+
inside = (ar >= starts.unsqueeze(-1)) & (ar < (starts + widths).unsqueeze(-1))
|
|
168
|
+
return inside.any(dim=1)
|
|
169
|
+
|
|
170
|
+
def forward(self, x: torch.Tensor, lengths: Optional[torch.Tensor] = None
|
|
171
|
+
) -> torch.Tensor:
|
|
172
|
+
c = self.cfg
|
|
173
|
+
if not self.training or not c.spec_augment:
|
|
174
|
+
return x
|
|
175
|
+
b, t, f = x.shape
|
|
176
|
+
dev = x.device
|
|
177
|
+
lens = (lengths.to(dev).clamp(1, t) if lengths is not None
|
|
178
|
+
else torch.full((b,), t, device=dev, dtype=torch.long))
|
|
179
|
+
if c.freq_masks:
|
|
180
|
+
fw = torch.full((b,), c.freq_mask_width, device=dev, dtype=torch.long)
|
|
181
|
+
x = x.masked_fill(self._bands(c.freq_masks, fw, f, b, dev).unsqueeze(1), 0.0)
|
|
182
|
+
if c.time_masks:
|
|
183
|
+
tw = (c.time_mask_pct * lens.float()).long().clamp(min=1)
|
|
184
|
+
tm = self._bands(c.time_masks, tw, t, b, dev)
|
|
185
|
+
valid = torch.arange(t, device=dev).view(1, t) < lens.view(b, 1)
|
|
186
|
+
x = x.masked_fill((tm & valid).unsqueeze(2), 0.0)
|
|
187
|
+
return x
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
# --------------------------------------------------------------------------- #
|
|
191
|
+
# Speed perturbation (Ko et al., 2015)
|
|
192
|
+
# --------------------------------------------------------------------------- #
|
|
193
|
+
def speed_perturb(wav: torch.Tensor, sample_rate: int,
|
|
194
|
+
factors=(0.9, 1.0, 1.1)) -> torch.Tensor:
|
|
195
|
+
"""Randomly retime a waveform by resampling: factor 0.9 = slower/longer.
|
|
196
|
+
|
|
197
|
+
Cheap "3x more data" that is very effective on small or noisy corpora
|
|
198
|
+
(different speaking rates, pitch shifts). Training-time only, raw-feature
|
|
199
|
+
mode only (precomputed mels are fixed at prep time).
|
|
200
|
+
"""
|
|
201
|
+
import random
|
|
202
|
+
factor = random.choice(list(factors))
|
|
203
|
+
if abs(factor - 1.0) < 1e-6:
|
|
204
|
+
return wav
|
|
205
|
+
import torchaudio.functional as AF
|
|
206
|
+
# Resampling to sr/factor then playing at sr changes duration by 1/factor.
|
|
207
|
+
return AF.resample(wav, orig_freq=int(sample_rate * factor),
|
|
208
|
+
new_freq=sample_rate)
|
altasr/benchmark.py
ADDED
|
@@ -0,0 +1,374 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
ALTASR :: benchmark CLI (installed as ``altasr-benchmark``).
|
|
4
|
+
|
|
5
|
+
Evaluate one or more ALTASR checkpoints -- and, optionally, OpenAI Whisper --
|
|
6
|
+
on the same labelled test set, and generate a professional, publishable
|
|
7
|
+
benchmark report (Markdown, and optionally HTML).
|
|
8
|
+
|
|
9
|
+
altasr-benchmark \\
|
|
10
|
+
--checkpoint "ALTASR small"=runs/small_full/best \\
|
|
11
|
+
--checkpoint "ALTASR medium"=runs/medium_full/best \\
|
|
12
|
+
--whisper small \\
|
|
13
|
+
--audio-root /data/track_b --metadata /data/track_b/dev_test.json \\
|
|
14
|
+
--limit 1000 --out benchmark/report.md --html
|
|
15
|
+
|
|
16
|
+
What the report contains
|
|
17
|
+
------------------------
|
|
18
|
+
* environment/methodology section (date, dataset, #utterances, hours,
|
|
19
|
+
hardware, decoding settings, normalisation policy);
|
|
20
|
+
* one row per system: parameters, WER, CER, punctuation-insensitive WER,
|
|
21
|
+
substitution/deletion/insertion rates, RTF (real-time factor);
|
|
22
|
+
* the hardest utterances per system (highest WER) as REF/HYP pairs;
|
|
23
|
+
* raw per-system JSON next to the report for further analysis/plotting.
|
|
24
|
+
|
|
25
|
+
Fairness notes baked in
|
|
26
|
+
-----------------------
|
|
27
|
+
* All hypotheses AND references are passed through ALTASR's normaliser before
|
|
28
|
+
scoring, so casing/punctuation conventions can't bias the comparison;
|
|
29
|
+
* ``wer_nopunct`` is reported for every system, so punctuated ALTASR models
|
|
30
|
+
are also compared on pure word accuracy;
|
|
31
|
+
* RTF is measured per system on the same machine in the same run.
|
|
32
|
+
|
|
33
|
+
Whisper support: ``pip install altasr[whisper]`` (faster-whisper). If the
|
|
34
|
+
package or its model download is unavailable, the benchmark simply proceeds
|
|
35
|
+
with the systems it can run and says so in the report.
|
|
36
|
+
"""
|
|
37
|
+
from __future__ import annotations
|
|
38
|
+
|
|
39
|
+
import argparse
|
|
40
|
+
import datetime as _dt
|
|
41
|
+
import json
|
|
42
|
+
import os
|
|
43
|
+
import platform
|
|
44
|
+
import sys
|
|
45
|
+
import time
|
|
46
|
+
import warnings
|
|
47
|
+
from typing import Callable, Dict, List, Optional, Tuple
|
|
48
|
+
|
|
49
|
+
warnings.filterwarnings("ignore", message=".*cuda capability.*")
|
|
50
|
+
warnings.filterwarnings("ignore", category=UserWarning, module="torch")
|
|
51
|
+
warnings.filterwarnings("ignore", category=FutureWarning)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
# --------------------------------------------------------------------------- #
|
|
55
|
+
# Systems
|
|
56
|
+
# --------------------------------------------------------------------------- #
|
|
57
|
+
class System:
|
|
58
|
+
"""name + transcribe(path)->text + metadata; one benchmark row."""
|
|
59
|
+
|
|
60
|
+
def __init__(self, name: str, transcribe: Callable[[str], str],
|
|
61
|
+
params_m: Optional[float], notes: str = "") -> None:
|
|
62
|
+
self.name = name
|
|
63
|
+
self.transcribe = transcribe
|
|
64
|
+
self.params_m = params_m
|
|
65
|
+
self.notes = notes
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _altasr_system(name: str, checkpoint: str, device: str,
|
|
69
|
+
quantize: str) -> System:
|
|
70
|
+
from altasr import Transcriber
|
|
71
|
+
asr = Transcriber.from_pretrained(checkpoint, device=device,
|
|
72
|
+
quantize=quantize)
|
|
73
|
+
notes = f"greedy CTC, device={asr.device}" + (
|
|
74
|
+
", int8" if quantize else "")
|
|
75
|
+
params = None
|
|
76
|
+
try:
|
|
77
|
+
params = asr.model.num_parameters() / 1e6
|
|
78
|
+
except Exception:
|
|
79
|
+
pass
|
|
80
|
+
return System(name, lambda p: asr.transcribe(p), params, notes)
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def _whisper_system(size: str, device: str) -> System:
|
|
84
|
+
"""faster-whisper preferred (CTranslate2, fair CPU speed); fall back to
|
|
85
|
+
openai-whisper if that's what is installed."""
|
|
86
|
+
try:
|
|
87
|
+
from faster_whisper import WhisperModel
|
|
88
|
+
dev = "cuda" if str(device).startswith("cuda") else "cpu"
|
|
89
|
+
model = WhisperModel(size, device=dev,
|
|
90
|
+
compute_type="float16" if dev == "cuda" else "int8")
|
|
91
|
+
|
|
92
|
+
def transcribe(path: str) -> str:
|
|
93
|
+
segments, _ = model.transcribe(path, language="rw",
|
|
94
|
+
beam_size=1, vad_filter=False)
|
|
95
|
+
return " ".join(s.text for s in segments)
|
|
96
|
+
|
|
97
|
+
return System(f"Whisper {size} (faster-whisper)", transcribe,
|
|
98
|
+
_WHISPER_PARAMS.get(size), f"beam=1, lang=rw, {dev}")
|
|
99
|
+
except ImportError:
|
|
100
|
+
pass
|
|
101
|
+
import whisper # raises ImportError with a clear message if absent
|
|
102
|
+
model = whisper.load_model(size, device="cuda" if str(device).startswith("cuda") else "cpu")
|
|
103
|
+
|
|
104
|
+
def transcribe(path: str) -> str:
|
|
105
|
+
return model.transcribe(path, language="rw", beam_size=1)["text"]
|
|
106
|
+
|
|
107
|
+
return System(f"Whisper {size}", transcribe,
|
|
108
|
+
_WHISPER_PARAMS.get(size), "beam=1, lang=rw")
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
_WHISPER_PARAMS = {"tiny": 39, "base": 74, "small": 244, "medium": 769,
|
|
112
|
+
"large": 1550, "large-v2": 1550, "large-v3": 1550}
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
# --------------------------------------------------------------------------- #
|
|
116
|
+
# Report generation
|
|
117
|
+
# --------------------------------------------------------------------------- #
|
|
118
|
+
def _md_table(rows: List[List[str]], header: List[str]) -> str:
|
|
119
|
+
out = ["| " + " | ".join(header) + " |",
|
|
120
|
+
"|" + "|".join("---" for _ in header) + "|"]
|
|
121
|
+
out += ["| " + " | ".join(r) + " |" for r in rows]
|
|
122
|
+
return "\n".join(out)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def write_report(path: str, dataset: Dict, results: List[Dict],
|
|
126
|
+
examples: Dict[str, List[Tuple[str, str, float]]],
|
|
127
|
+
skipped: List[str], html: bool) -> None:
|
|
128
|
+
now = _dt.datetime.now().strftime("%Y-%m-%d %H:%M")
|
|
129
|
+
hw = f"{platform.processor() or platform.machine()}, {platform.system()}"
|
|
130
|
+
try:
|
|
131
|
+
import torch
|
|
132
|
+
if torch.cuda.is_available():
|
|
133
|
+
hw += f", GPU: {torch.cuda.get_device_name(0)}"
|
|
134
|
+
except Exception:
|
|
135
|
+
pass
|
|
136
|
+
|
|
137
|
+
lines = [
|
|
138
|
+
f"# ASR Benchmark Report",
|
|
139
|
+
"",
|
|
140
|
+
f"*Generated by `altasr-benchmark` on {now}.*",
|
|
141
|
+
"",
|
|
142
|
+
"## Methodology",
|
|
143
|
+
"",
|
|
144
|
+
f"- **Test set:** {dataset['metadata']} "
|
|
145
|
+
f"({dataset['utts']} utterances, {dataset['hours']:.2f} h of audio)",
|
|
146
|
+
f"- **Hardware:** {hw}",
|
|
147
|
+
f"- **Scoring:** references and hypotheses are both normalised with "
|
|
148
|
+
f"ALTASR's text normaliser (lowercase, apostrophes folded, "
|
|
149
|
+
f"punctuation `{dataset['punctuation']}` kept and re-spaced) before "
|
|
150
|
+
f"computing corpus-level WER/CER. `WER (no punct)` strips punctuation "
|
|
151
|
+
f"from both sides, making systems that do and don't produce "
|
|
152
|
+
f"punctuation directly comparable.",
|
|
153
|
+
f"- **RTF** = processing seconds / audio seconds on this machine "
|
|
154
|
+
f"(lower is better; < 1.0 is faster than real time). All systems ran "
|
|
155
|
+
f"in the same process, one utterance at a time (batch size 1).",
|
|
156
|
+
"",
|
|
157
|
+
"## Results",
|
|
158
|
+
"",
|
|
159
|
+
]
|
|
160
|
+
header = ["System", "Params (M)", "WER %", "CER %", "WER no-punct %",
|
|
161
|
+
"Sub %", "Del %", "Ins %", "RTF"]
|
|
162
|
+
rows = []
|
|
163
|
+
for r in results:
|
|
164
|
+
m = r["metrics"]
|
|
165
|
+
rows.append([
|
|
166
|
+
f"**{r['name']}**",
|
|
167
|
+
f"{r['params_m']:.0f}" if r["params_m"] else "–",
|
|
168
|
+
f"{m['wer'] * 100:.2f}", f"{m['cer'] * 100:.2f}",
|
|
169
|
+
f"{m['wer_nopunct'] * 100:.2f}",
|
|
170
|
+
f"{m['sub_rate'] * 100:.2f}", f"{m['del_rate'] * 100:.2f}",
|
|
171
|
+
f"{m['ins_rate'] * 100:.2f}", f"{m['rtf']:.3f}"])
|
|
172
|
+
lines.append(_md_table(rows, header))
|
|
173
|
+
lines += ["", "### System notes", ""]
|
|
174
|
+
lines += [f"- **{r['name']}** — {r['notes']}" for r in results]
|
|
175
|
+
if skipped:
|
|
176
|
+
lines += ["", "### Skipped systems", ""]
|
|
177
|
+
lines += [f"- {s}" for s in skipped]
|
|
178
|
+
|
|
179
|
+
lines += ["", "## Hardest utterances per system", "",
|
|
180
|
+
"The five highest-WER validation utterances for each system — "
|
|
181
|
+
"useful for spotting systematic failure modes (noise, "
|
|
182
|
+
"code-switching, long-form audio, numerals).", ""]
|
|
183
|
+
for name, exs in examples.items():
|
|
184
|
+
lines += [f"### {name}", ""]
|
|
185
|
+
for ref, hyp, w in exs:
|
|
186
|
+
lines += [f"- WER {w * 100:.0f}%",
|
|
187
|
+
f" - REF: `{ref[:160]}`",
|
|
188
|
+
f" - HYP: `{hyp[:160] or '(empty)'}`"]
|
|
189
|
+
lines.append("")
|
|
190
|
+
|
|
191
|
+
md = "\n".join(lines) + "\n"
|
|
192
|
+
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
|
|
193
|
+
with open(path, "w", encoding="utf-8") as fh:
|
|
194
|
+
fh.write(md)
|
|
195
|
+
print(f"[report] {path}")
|
|
196
|
+
|
|
197
|
+
if html:
|
|
198
|
+
html_path = os.path.splitext(path)[0] + ".html"
|
|
199
|
+
with open(html_path, "w", encoding="utf-8") as fh:
|
|
200
|
+
fh.write(_md_to_html(md))
|
|
201
|
+
print(f"[report] {html_path}")
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
def _md_to_html(md: str) -> str:
|
|
205
|
+
"""A small dependency-free Markdown->HTML converter (headers, tables,
|
|
206
|
+
lists, bold, inline code) good enough for the benchmark report."""
|
|
207
|
+
import html as _h
|
|
208
|
+
import re
|
|
209
|
+
out, in_table, in_list = [], False, False
|
|
210
|
+
|
|
211
|
+
def inline(s: str) -> str:
|
|
212
|
+
s = _h.escape(s)
|
|
213
|
+
s = re.sub(r"\*\*(.+?)\*\*", r"<b>\1</b>", s)
|
|
214
|
+
s = re.sub(r"\*(.+?)\*", r"<i>\1</i>", s)
|
|
215
|
+
s = re.sub(r"`(.+?)`", r"<code>\1</code>", s)
|
|
216
|
+
return s
|
|
217
|
+
|
|
218
|
+
for line in md.splitlines():
|
|
219
|
+
if line.startswith("|"):
|
|
220
|
+
cells = [c.strip() for c in line.strip("|").split("|")]
|
|
221
|
+
if all(set(c) <= {"-", " "} for c in cells):
|
|
222
|
+
continue
|
|
223
|
+
if not in_table:
|
|
224
|
+
out.append("<table>"); in_table = True
|
|
225
|
+
out.append("<tr>" + "".join(f"<th>{inline(c)}</th>"
|
|
226
|
+
for c in cells) + "</tr>")
|
|
227
|
+
else:
|
|
228
|
+
out.append("<tr>" + "".join(f"<td>{inline(c)}</td>"
|
|
229
|
+
for c in cells) + "</tr>")
|
|
230
|
+
continue
|
|
231
|
+
if in_table:
|
|
232
|
+
out.append("</table>"); in_table = False
|
|
233
|
+
if line.startswith("- ") or line.startswith(" - "):
|
|
234
|
+
if not in_list:
|
|
235
|
+
out.append("<ul>"); in_list = True
|
|
236
|
+
out.append(f"<li>{inline(line.lstrip(' -'))}</li>")
|
|
237
|
+
continue
|
|
238
|
+
if in_list:
|
|
239
|
+
out.append("</ul>"); in_list = False
|
|
240
|
+
m = re.match(r"^(#{1,3}) (.*)$", line)
|
|
241
|
+
if m:
|
|
242
|
+
n = len(m.group(1))
|
|
243
|
+
out.append(f"<h{n}>{inline(m.group(2))}</h{n}>")
|
|
244
|
+
elif line.strip():
|
|
245
|
+
out.append(f"<p>{inline(line)}</p>")
|
|
246
|
+
if in_table:
|
|
247
|
+
out.append("</table>")
|
|
248
|
+
if in_list:
|
|
249
|
+
out.append("</ul>")
|
|
250
|
+
style = ("<style>body{font-family:system-ui,sans-serif;max-width:960px;"
|
|
251
|
+
"margin:2rem auto;padding:0 1rem;line-height:1.5}"
|
|
252
|
+
"table{border-collapse:collapse;width:100%}"
|
|
253
|
+
"th,td{border:1px solid #ccc;padding:6px 10px;text-align:left}"
|
|
254
|
+
"th{background:#f3f3f3}code{background:#f5f5f5;padding:1px 4px;"
|
|
255
|
+
"border-radius:3px}</style>")
|
|
256
|
+
return ("<!doctype html><html><head><meta charset='utf-8'>"
|
|
257
|
+
f"<title>ASR Benchmark Report</title>{style}</head><body>"
|
|
258
|
+
+ "\n".join(out) + "</body></html>")
|
|
259
|
+
|
|
260
|
+
|
|
261
|
+
# --------------------------------------------------------------------------- #
|
|
262
|
+
# Main
|
|
263
|
+
# --------------------------------------------------------------------------- #
|
|
264
|
+
def main(argv=None) -> int:
|
|
265
|
+
from altasr.utils import clear_terminal
|
|
266
|
+
clear_terminal()
|
|
267
|
+
p = argparse.ArgumentParser(
|
|
268
|
+
description="Benchmark ALTASR checkpoints (and optionally Whisper) "
|
|
269
|
+
"and generate a report.",
|
|
270
|
+
formatter_class=argparse.ArgumentDefaultsHelpFormatter)
|
|
271
|
+
p.add_argument("--checkpoint", action="append", required=True,
|
|
272
|
+
metavar="[NAME=]PATH",
|
|
273
|
+
help="ALTASR checkpoint folder; repeatable. Optional "
|
|
274
|
+
"'Name=' prefix for the report row.")
|
|
275
|
+
p.add_argument("--whisper", action="append", default=[],
|
|
276
|
+
metavar="SIZE", help="Also benchmark Whisper (tiny/base/"
|
|
277
|
+
"small/medium/large-v3); needs altasr[whisper]. Repeatable.")
|
|
278
|
+
p.add_argument("--audio-root", required=True)
|
|
279
|
+
p.add_argument("--metadata", nargs="+", required=True)
|
|
280
|
+
p.add_argument("--limit", type=int, default=0,
|
|
281
|
+
help="Evaluate only the first N utterances (0 = all).")
|
|
282
|
+
p.add_argument("--out", default="benchmark/report.md")
|
|
283
|
+
p.add_argument("--html", action="store_true",
|
|
284
|
+
help="Also write an HTML version of the report.")
|
|
285
|
+
p.add_argument("--device", default="auto",
|
|
286
|
+
help="auto | cpu | cuda | cuda:N")
|
|
287
|
+
p.add_argument("--quantize", choices=["", "int8"], default="",
|
|
288
|
+
help="Apply int8 quantization to the ALTASR systems "
|
|
289
|
+
"(CPU edge-deployment benchmark).")
|
|
290
|
+
args = p.parse_args(argv)
|
|
291
|
+
|
|
292
|
+
from altasr.config import Config
|
|
293
|
+
from altasr.metrics import corpus_metrics, wer as _wer
|
|
294
|
+
from altasr.text import DEFAULT_PUNCTUATION, normalize_text
|
|
295
|
+
from altasr.training.data import load_records
|
|
296
|
+
|
|
297
|
+
cfg = Config()
|
|
298
|
+
cfg.data.audio_root = args.audio_root
|
|
299
|
+
records = load_records(cfg, args.metadata, check_files=True)
|
|
300
|
+
if args.limit:
|
|
301
|
+
records = records[: args.limit]
|
|
302
|
+
if not records:
|
|
303
|
+
print("ERROR: no records resolved (check --audio-root/--metadata).",
|
|
304
|
+
file=sys.stderr)
|
|
305
|
+
return 2
|
|
306
|
+
hours = sum(r["duration"] or 0.0 for r in records) / 3600.0
|
|
307
|
+
punct = cfg.text.punctuation
|
|
308
|
+
print(f"[bench] {len(records)} utterances (~{hours:.2f} h)", flush=True)
|
|
309
|
+
|
|
310
|
+
# -- build systems -------------------------------------------------------- #
|
|
311
|
+
systems: List[System] = []
|
|
312
|
+
skipped: List[str] = []
|
|
313
|
+
for spec in args.checkpoint:
|
|
314
|
+
name, _, path = spec.rpartition("=")
|
|
315
|
+
name = name or f"ALTASR ({os.path.basename(os.path.dirname(path.rstrip('/'))) or path})"
|
|
316
|
+
try:
|
|
317
|
+
systems.append(_altasr_system(name, path, args.device, args.quantize))
|
|
318
|
+
except Exception as exc:
|
|
319
|
+
skipped.append(f"{name}: failed to load ({exc})")
|
|
320
|
+
for size in args.whisper:
|
|
321
|
+
try:
|
|
322
|
+
systems.append(_whisper_system(size, args.device))
|
|
323
|
+
except Exception as exc:
|
|
324
|
+
skipped.append(f"Whisper {size}: unavailable ({exc}) — "
|
|
325
|
+
f"`pip install altasr[whisper]`")
|
|
326
|
+
if not systems:
|
|
327
|
+
print("ERROR: no system could be loaded.", file=sys.stderr)
|
|
328
|
+
for s in skipped:
|
|
329
|
+
print(" -", s, file=sys.stderr)
|
|
330
|
+
return 2
|
|
331
|
+
|
|
332
|
+
# -- run -------------------------------------------------------------------- #
|
|
333
|
+
results, examples = [], {}
|
|
334
|
+
for system in systems:
|
|
335
|
+
print(f"[bench] running {system.name} ...", flush=True)
|
|
336
|
+
pairs: List[Tuple[str, str]] = []
|
|
337
|
+
t0 = time.time()
|
|
338
|
+
for i, rec in enumerate(records, 1):
|
|
339
|
+
try:
|
|
340
|
+
hyp = system.transcribe(rec["resolved"])
|
|
341
|
+
except Exception as exc:
|
|
342
|
+
hyp = ""
|
|
343
|
+
print(f"\n [warn] {rec['id']}: {exc}", flush=True)
|
|
344
|
+
hyp = normalize_text(hyp, keep_punctuation=True, punctuation=punct)
|
|
345
|
+
pairs.append((rec["text"], hyp))
|
|
346
|
+
if i % 20 == 0 or i == len(records):
|
|
347
|
+
print(f"\r {i}/{len(records)}", end="", flush=True)
|
|
348
|
+
elapsed = time.time() - t0
|
|
349
|
+
print()
|
|
350
|
+
m = corpus_metrics(pairs, punctuation=punct,
|
|
351
|
+
audio_seconds=hours * 3600.0,
|
|
352
|
+
processing_seconds=elapsed)
|
|
353
|
+
results.append({"name": system.name, "params_m": system.params_m,
|
|
354
|
+
"notes": system.notes, "metrics": m})
|
|
355
|
+
worst = sorted(((r, h, _wer(r, h)) for r, h in pairs),
|
|
356
|
+
key=lambda x: -x[2])[:5]
|
|
357
|
+
examples[system.name] = worst
|
|
358
|
+
print(f" WER {m['wer'] * 100:.2f}% CER {m['cer'] * 100:.2f}% "
|
|
359
|
+
f"WER(np) {m['wer_nopunct'] * 100:.2f}% RTF {m['rtf']:.3f}",
|
|
360
|
+
flush=True)
|
|
361
|
+
|
|
362
|
+
dataset = {"metadata": ", ".join(args.metadata), "utts": len(records),
|
|
363
|
+
"hours": hours, "punctuation": punct or DEFAULT_PUNCTUATION}
|
|
364
|
+
write_report(args.out, dataset, results, examples, skipped, args.html)
|
|
365
|
+
with open(os.path.splitext(args.out)[0] + ".json", "w",
|
|
366
|
+
encoding="utf-8") as fh:
|
|
367
|
+
json.dump({"dataset": dataset,
|
|
368
|
+
"results": [{k: v for k, v in r.items()} for r in results]},
|
|
369
|
+
fh, indent=2)
|
|
370
|
+
return 0
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
if __name__ == "__main__":
|
|
374
|
+
raise SystemExit(main())
|