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