|
| 1 | +"""`assembly eval` — transcribe an evaluation dataset and score it against references. |
| 2 | +
|
| 3 | +WER (via jiwer) against the dataset's reference texts; with ``--speaker-labels`` |
| 4 | +also DER (via pyannote.metrics) against its reference speaker turns. The module |
| 5 | +is named ``evaluate`` because importing a module named ``eval`` would shadow the |
| 6 | +builtin; the command itself registers as ``eval``. |
| 7 | +""" |
| 8 | + |
| 9 | +from __future__ import annotations |
| 10 | + |
| 11 | +from dataclasses import dataclass |
| 12 | + |
| 13 | +import assemblyai as aai |
| 14 | +import typer |
| 15 | +from rich.console import RenderableType |
| 16 | + |
| 17 | +from aai_cli import client, config, der, eval_data, help_panels, jsonshape, options, output, wer |
| 18 | +from aai_cli.context import AppState, run_command |
| 19 | +from aai_cli.help_text import examples_epilog |
| 20 | + |
| 21 | +app = typer.Typer() |
| 22 | + |
| 23 | + |
| 24 | +def _pct(value: object) -> str: |
| 25 | + return f"{jsonshape.as_float(value):.2%}" |
| 26 | + |
| 27 | + |
| 28 | +def _hypothesis_turns(transcript: aai.Transcript) -> list[der.Turn]: |
| 29 | + """The transcript's diarized utterances as DER hypothesis turns (ms → seconds).""" |
| 30 | + return [ |
| 31 | + der.Turn( |
| 32 | + speaker=str(getattr(utterance, "speaker", "")), |
| 33 | + start=jsonshape.as_float(getattr(utterance, "start", None)) / 1000, |
| 34 | + end=jsonshape.as_float(getattr(utterance, "end", None)) / 1000, |
| 35 | + ) |
| 36 | + for utterance in jsonshape.object_list(getattr(transcript, "utterances", None)) |
| 37 | + ] |
| 38 | + |
| 39 | + |
| 40 | +@dataclass(frozen=True) |
| 41 | +class _ItemResult: |
| 42 | + """One scored row: the emitted dict plus the scores kept for pooling.""" |
| 43 | + |
| 44 | + row: dict[str, object] |
| 45 | + words: wer.Score | None |
| 46 | + speakers: der.DerScore | None |
| 47 | + |
| 48 | + |
| 49 | +def _score_item( |
| 50 | + item: eval_data.EvalItem, transcript: aai.Transcript, *, collar: float |
| 51 | +) -> _ItemResult: |
| 52 | + row: dict[str, object] = {"item": item.item_id} |
| 53 | + words: wer.Score | None = None |
| 54 | + speakers: der.DerScore | None = None |
| 55 | + if item.reference is not None: |
| 56 | + words = wer.score(item.reference, str(transcript.text or "")) |
| 57 | + row.update({"words": words.words, "errors": words.errors, "wer": words.wer}) |
| 58 | + if item.turns is not None: |
| 59 | + speakers = der.score(item.turns, _hypothesis_turns(transcript), collar=collar) |
| 60 | + row["der"] = speakers.der |
| 61 | + return _ItemResult(row=row, words=words, speakers=speakers) |
| 62 | + |
| 63 | + |
| 64 | +def _payload( |
| 65 | + label: str, speech_model: aai.SpeechModel | None, results: list[_ItemResult] |
| 66 | +) -> dict[str, object]: |
| 67 | + payload: dict[str, object] = { |
| 68 | + "dataset": label, |
| 69 | + "speech_model": speech_model.value if speech_model else None, |
| 70 | + "items": len(results), |
| 71 | + "rows": [result.row for result in results], |
| 72 | + } |
| 73 | + word_scores = [result.words for result in results if result.words is not None] |
| 74 | + if word_scores: |
| 75 | + total = wer.pooled(word_scores) |
| 76 | + payload.update({"words": total.words, "errors": total.errors, "wer": total.wer}) |
| 77 | + der_scores = [result.speakers for result in results if result.speakers is not None] |
| 78 | + if der_scores: |
| 79 | + payload["der"] = der.pooled(der_scores).der |
| 80 | + return payload |
| 81 | + |
| 82 | + |
| 83 | +def _summary(payload: dict[str, object]) -> str: |
| 84 | + parts: list[str] = [] |
| 85 | + if "wer" in payload: |
| 86 | + errors = jsonshape.as_int(payload.get("errors")) |
| 87 | + noun = "error" if errors == 1 else "errors" |
| 88 | + parts.append( |
| 89 | + f"WER {_pct(payload.get('wer'))} ({errors} {noun} / {payload.get('words')} words)" |
| 90 | + ) |
| 91 | + if "der" in payload: |
| 92 | + parts.append(f"DER {_pct(payload.get('der'))}") |
| 93 | + return output.heading(" ".join(parts)) |
| 94 | + |
| 95 | + |
| 96 | +def _render(payload: dict[str, object]) -> RenderableType: |
| 97 | + has_wer = "wer" in payload |
| 98 | + has_der = "der" in payload |
| 99 | + columns = [ |
| 100 | + "ITEM", |
| 101 | + *(["WORDS", "ERRORS", "WER"] if has_wer else []), |
| 102 | + *(["DER"] if has_der else []), |
| 103 | + ] |
| 104 | + table = output.data_table(*columns) |
| 105 | + for row in jsonshape.mapping_list(payload.get("rows")): |
| 106 | + cells = [str(row.get("item"))] |
| 107 | + if has_wer: |
| 108 | + cells += [str(row.get("words")), str(row.get("errors")), _pct(row.get("wer"))] |
| 109 | + if has_der: |
| 110 | + cells.append(_pct(row.get("der"))) |
| 111 | + table.add_row(*cells) |
| 112 | + model = payload.get("speech_model") or "default model" |
| 113 | + return output.stack( |
| 114 | + output.muted(f"{payload.get('dataset')} · {model}"), table, _summary(payload) |
| 115 | + ) |
| 116 | + |
| 117 | + |
| 118 | +@app.command( |
| 119 | + name="eval", |
| 120 | + rich_help_panel=help_panels.TRANSCRIPTION, |
| 121 | + epilog=examples_epilog( |
| 122 | + [ |
| 123 | + ("Score a model on 10 rows of an HF dataset", "assembly eval distil-whisper/meanwhile"), |
| 124 | + ( |
| 125 | + "Compare models on your own audio", |
| 126 | + "assembly eval calls.csv --speech-model universal", |
| 127 | + ), |
| 128 | + ( |
| 129 | + "Score diarization too (WER + DER)", |
| 130 | + "assembly eval agent-calls.jsonl --speaker-labels", |
| 131 | + ), |
| 132 | + ( |
| 133 | + "Pick a subset/split and more rows", |
| 134 | + "assembly eval mozilla-foundation/common_voice_17_0 --subset en --limit 50", |
| 135 | + ), |
| 136 | + ] |
| 137 | + ), |
| 138 | +) |
| 139 | +def evaluate( |
| 140 | + ctx: typer.Context, |
| 141 | + dataset: str = typer.Argument( |
| 142 | + ..., |
| 143 | + help="Hugging Face dataset id, or a local .csv/.jsonl manifest with audio + text columns.", |
| 144 | + ), |
| 145 | + split: str | None = typer.Option( |
| 146 | + None, "--split", help="Hugging Face split to score (default: test)." |
| 147 | + ), |
| 148 | + subset: str | None = typer.Option( |
| 149 | + None, "--subset", help="Hugging Face config/subset name (e.g. a language)." |
| 150 | + ), |
| 151 | + limit: int = typer.Option(10, "--limit", min=1, max=100, help="Rows to evaluate (1-100)."), |
| 152 | + audio_column: str | None = typer.Option( |
| 153 | + None, "--audio-column", help="Audio column name (default: auto-detect)." |
| 154 | + ), |
| 155 | + text_column: str | None = typer.Option( |
| 156 | + None, "--text-column", help="Reference text column name (default: auto-detect)." |
| 157 | + ), |
| 158 | + speech_model: aai.SpeechModel | None = typer.Option( |
| 159 | + None, "--speech-model", help="Speech model to evaluate." |
| 160 | + ), |
| 161 | + language_code: str | None = typer.Option( |
| 162 | + None, "--language-code", help="Force a language (e.g. en_us)." |
| 163 | + ), |
| 164 | + speaker_labels: bool = typer.Option( |
| 165 | + False, |
| 166 | + "--speaker-labels", |
| 167 | + help="Diarize and also score DER against the dataset's reference speaker turns (speakers/timestamps_start/timestamps_end columns, in seconds).", |
| 168 | + ), |
| 169 | + collar: float = typer.Option( |
| 170 | + 1.0, |
| 171 | + "--collar", |
| 172 | + min=0.0, |
| 173 | + help="DER forgiveness (seconds) around each reference turn boundary.", |
| 174 | + ), |
| 175 | + json_out: bool = options.json_option("Output the rows and summary as one JSON object."), |
| 176 | +) -> None: |
| 177 | + """Transcribe an evaluation dataset and score WER against its reference texts. |
| 178 | +
|
| 179 | + Handy for picking a model: run once per --speech-model and compare. Datasets |
| 180 | + come from the Hugging Face Hub (gated ones need HF_TOKEN) or a local .csv/.jsonl |
| 181 | + manifest. --speaker-labels also scores diarization (DER) against reference |
| 182 | + speaker turns. |
| 183 | + """ |
| 184 | + |
| 185 | + def body(state: AppState, json_mode: bool) -> None: |
| 186 | + data = eval_data.load( |
| 187 | + dataset, |
| 188 | + split=split, |
| 189 | + subset=subset, |
| 190 | + audio_column=audio_column, |
| 191 | + text_column=text_column, |
| 192 | + limit=limit, |
| 193 | + with_speakers=speaker_labels, |
| 194 | + ) |
| 195 | + api_key = config.resolve_api_key(profile=state.profile) |
| 196 | + transcription_config = aai.TranscriptionConfig( |
| 197 | + speech_model=speech_model, |
| 198 | + language_code=language_code, |
| 199 | + speaker_labels=speaker_labels or None, |
| 200 | + ) |
| 201 | + results: list[_ItemResult] = [] |
| 202 | + for index, item in enumerate(data.items, start=1): |
| 203 | + with output.status( |
| 204 | + f"[{index}/{len(data.items)}] Transcribing {item.item_id}…", |
| 205 | + json_mode=json_mode, |
| 206 | + quiet=state.quiet, |
| 207 | + ): |
| 208 | + transcript = client.transcribe(api_key, item.audio, config=transcription_config) |
| 209 | + results.append(_score_item(item, transcript, collar=collar)) |
| 210 | + output.emit(_payload(data.label, speech_model, results), _render, json_mode=json_mode) |
| 211 | + |
| 212 | + run_command(ctx, body, json=json_out) |
0 commit comments