Skip to content

Commit 73a1dec

Browse files
alexkromanclaude
andauthored
Add assembly eval: dataset WER/DER scoring for model selection (#74)
Transcribe an evaluation dataset — a Hugging Face dataset id (via the datasets-server REST API, no datasets dependency) or a local .csv/.jsonl manifest — and score word error rate against its reference texts with jiwer. With --speaker-labels the run also diarizes and scores diarization error rate against reference speaker turns via pyannote.metrics (--collar for boundary forgiveness). Per-file rows plus pooled corpus scores render as a table or --json. https://claude.ai/code/session_014o9JsqBmqmybikbtHNsbSf Co-authored-by: Claude <noreply@anthropic.com>
1 parent 246348f commit 73a1dec

16 files changed

Lines changed: 2310 additions & 8 deletions

‎.importlinter‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,8 +14,10 @@ source_modules =
1414
aai_cli.config
1515
aai_cli.config_builder
1616
aai_cli.context
17+
aai_cli.der
1718
aai_cli.environments
1819
aai_cli.errors
20+
aai_cli.eval_data
1921
aai_cli.follow
2022
aai_cli.help_panels
2123
aai_cli.help_text
@@ -32,6 +34,7 @@ source_modules =
3234
aai_cli.transcribe_batch
3335
aai_cli.transcribe_exec
3436
aai_cli.transcribe_render
37+
aai_cli.wer
3538
aai_cli.ws
3639
aai_cli.youtube
3740
forbidden_modules =
@@ -47,6 +50,7 @@ modules =
4750
aai_cli.commands.deploy
4851
aai_cli.commands.dev
4952
aai_cli.commands.doctor
53+
aai_cli.commands.evaluate
5054
aai_cli.commands.init
5155
aai_cli.commands.keys
5256
aai_cli.commands.llm
@@ -66,9 +70,12 @@ source_modules =
6670
aai_cli.client
6771
aai_cli.config
6872
aai_cli.config_builder
73+
aai_cli.der
6974
aai_cli.environments
7075
aai_cli.errors
76+
aai_cli.eval_data
7177
aai_cli.llm
7278
aai_cli.telemetry
79+
aai_cli.wer
7380
forbidden_modules =
7481
rich

‎README.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ Requires Python 3.12+. On Linux, install PortAudio once for microphone support (
5050
- **Real-time streaming**: `assembly stream` transcribes the microphone, a file, or a URL live — on macOS it can capture system audio too.
5151
- **Voice agent**: `assembly agent` runs a full-duplex spoken conversation in your terminal (use headphones).
5252
- **LLM Gateway**: `assembly llm` prompts an LLM over a transcript, stdin, or a live stream (`assembly stream --llm "summarize as I talk"`).
53+
- **Model evaluation**: `assembly eval` transcribes a Hugging Face dataset or a local `.csv`/`.jsonl` manifest and scores WER against its references (plus DER with `--speaker-labels`) — handy for picking a speech model.
5354
- **Starter apps**: `assembly init` scaffolds a self-contained FastAPI + HTML app (`audio-transcription`, `live-captions`, `voice-agent`).
5455
- **Code generation**: add `--show-code` to `transcribe`/`stream`/`agent` to print the equivalent Python SDK script instead of running.
5556
- **Account self-service**: `assembly keys` / `balance` / `usage` / `limits` / `sessions` / `audit` via browser login.

‎aai_cli/commands/evaluate.py‎

Lines changed: 212 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,212 @@
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)

‎aai_cli/der.py‎

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,95 @@
1+
"""Diarization error rate (DER) scoring for `assembly eval --speaker-labels`.
2+
3+
A thin shim over :mod:`pyannote.metrics` — the de-facto standard DER
4+
implementation, including its optimal reference↔hypothesis speaker mapping —
5+
so the alignment math is never re-derived here. pyannote drags numpy/scipy/
6+
pandas along, so it is imported lazily inside :func:`score`; the dataclasses
7+
here stay import-cheap for the command layer. No SDK, no Rich.
8+
"""
9+
10+
from __future__ import annotations
11+
12+
from dataclasses import dataclass
13+
from typing import TYPE_CHECKING
14+
15+
from pydantic import TypeAdapter
16+
17+
if TYPE_CHECKING:
18+
from pyannote.core import Annotation
19+
20+
# pyannote types the metric call as a bare float; with ``detailed=True`` it
21+
# actually returns the components dict — validate that shape instead of casting.
22+
_COMPONENTS: TypeAdapter[dict[str, float]] = TypeAdapter(dict[str, float])
23+
24+
25+
@dataclass(frozen=True)
26+
class Turn:
27+
"""One speaker turn: who spoke from ``start`` to ``end`` (seconds)."""
28+
29+
speaker: str
30+
start: float
31+
end: float
32+
33+
34+
@dataclass(frozen=True)
35+
class DerScore:
36+
"""DER components in seconds of speech; pooled across files for corpus DER."""
37+
38+
missed: float
39+
false_alarm: float
40+
confusion: float
41+
total: float
42+
43+
@property
44+
def der(self) -> float:
45+
return (self.missed + self.false_alarm + self.confusion) / self.total
46+
47+
48+
def _annotation(turns: list[Turn]) -> Annotation:
49+
from pyannote.core import Annotation, Segment
50+
51+
annotation = Annotation()
52+
# The list index keys each track, so two overlapping turns by the same
53+
# speaker don't collapse into one.
54+
for track, turn in enumerate(turns):
55+
annotation[Segment(turn.start, turn.end), track] = turn.speaker
56+
return annotation
57+
58+
59+
def score(reference: list[Turn], hypothesis: list[Turn], *, collar: float = 0.0) -> DerScore:
60+
"""Score hypothesis speaker turns against the reference diarization.
61+
62+
The caller guarantees a non-empty reference (the dataset loader rejects
63+
rows without speaker turns), so ``DerScore.der`` is always well-defined.
64+
``collar`` forgives that many seconds around each reference turn boundary.
65+
"""
66+
from pyannote.core import Segment, Timeline
67+
from pyannote.metrics.diarization import DiarizationErrorRate
68+
69+
metric = DiarizationErrorRate(collar=collar)
70+
# An explicit evaluation extent (audio start through the last turn either
71+
# side heard); without it pyannote warns about approximating the UEM.
72+
extent = max(turn.end for turn in [*reference, *hypothesis])
73+
components: object = metric(
74+
_annotation(reference),
75+
_annotation(hypothesis),
76+
uem=Timeline([Segment(0.0, extent)]),
77+
detailed=True,
78+
)
79+
detail = _COMPONENTS.validate_python(components)
80+
return DerScore(
81+
missed=detail["missed detection"],
82+
false_alarm=detail["false alarm"],
83+
confusion=detail["confusion"],
84+
total=detail["total"],
85+
)
86+
87+
88+
def pooled(scores: list[DerScore]) -> DerScore:
89+
"""Corpus-level score: error seconds over total reference speech seconds."""
90+
return DerScore(
91+
missed=sum(item.missed for item in scores),
92+
false_alarm=sum(item.false_alarm for item in scores),
93+
confusion=sum(item.confusion for item in scores),
94+
total=sum(item.total for item in scores),
95+
)

0 commit comments

Comments
 (0)