Skip to content

Commit b493270

Browse files
committed
Extend eval --speech-model to all SDK-supported models
The SDK's SpeechModel enum stops at the legacy generation (best, nano, slam-1, universal); the current models (universal-3-pro, universal-2) are requested through its newer speech_models list parameter. Replace the option's type with a CLI-level enum covering the union and route each choice through the SDK parameter that accepts it. https://claude.ai/code/session_01Ra3Q7AeGdx6uFxTtc7dL1K
1 parent 66d3636 commit b493270

3 files changed

Lines changed: 75 additions & 10 deletions

File tree

‎aai_cli/commands/evaluate.py‎

Lines changed: 47 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from __future__ import annotations
1010

1111
from dataclasses import dataclass
12+
from enum import StrEnum
1213

1314
import assemblyai as aai
1415
import typer
@@ -21,6 +22,47 @@
2122
app = typer.Typer()
2223

2324

25+
class EvalSpeechModel(StrEnum):
26+
"""Every speech model the SDK can request, current generation first.
27+
28+
The SDK's own ``SpeechModel`` enum stops at the legacy generation; the current
29+
models are requested through its newer ``speech_models`` list parameter, so the
30+
CLI choices are the union of both.
31+
"""
32+
33+
universal_3_pro = "universal-3-pro"
34+
universal_2 = "universal-2"
35+
slam_1 = "slam-1"
36+
universal = "universal"
37+
best = "best"
38+
nano = "nano"
39+
40+
41+
def _transcription_config(
42+
speech_model: EvalSpeechModel | None, *, language_code: str | None, speaker_labels: bool
43+
) -> aai.TranscriptionConfig:
44+
"""Route the model choice through the SDK parameter that accepts it.
45+
46+
Legacy models go through ``speech_model`` (the SDK enum); current-generation
47+
models only exist as ``speech_models`` list values.
48+
"""
49+
legacy = {model.value for model in aai.SpeechModel}
50+
return aai.TranscriptionConfig(
51+
speech_model=(
52+
aai.SpeechModel(speech_model.value)
53+
if speech_model is not None and speech_model.value in legacy
54+
else None
55+
),
56+
speech_models=(
57+
[speech_model.value]
58+
if speech_model is not None and speech_model.value not in legacy
59+
else None
60+
),
61+
language_code=language_code,
62+
speaker_labels=speaker_labels or None,
63+
)
64+
65+
2466
def _pct(value: object) -> str:
2567
return f"{jsonshape.as_float(value):.2%}"
2668

@@ -62,7 +104,7 @@ def _score_item(
62104

63105

64106
def _payload(
65-
label: str, speech_model: aai.SpeechModel | None, results: list[_ItemResult]
107+
label: str, speech_model: EvalSpeechModel | None, results: list[_ItemResult]
66108
) -> dict[str, object]:
67109
payload: dict[str, object] = {
68110
"dataset": label,
@@ -123,7 +165,7 @@ def _render(payload: dict[str, object]) -> RenderableType:
123165
("Score a model on 10 rows of an HF dataset", "assembly eval distil-whisper/meanwhile"),
124166
(
125167
"Compare models on your own audio",
126-
"assembly eval calls.csv --speech-model universal",
168+
"assembly eval calls.csv --speech-model universal-3-pro",
127169
),
128170
(
129171
"Score diarization too (WER + DER)",
@@ -163,7 +205,7 @@ def evaluate(
163205
text_column: str | None = typer.Option(
164206
None, "--text-column", help="Reference text column name (default: auto-detect)."
165207
),
166-
speech_model: aai.SpeechModel | None = typer.Option(
208+
speech_model: EvalSpeechModel | None = typer.Option(
167209
None, "--speech-model", help="Speech model to evaluate."
168210
),
169211
language_code: str | None = typer.Option(
@@ -212,10 +254,8 @@ def body(state: AppState, json_mode: bool) -> None:
212254
with_speakers=speaker_labels,
213255
)
214256
api_key = config.resolve_api_key(profile=state.profile)
215-
transcription_config = aai.TranscriptionConfig(
216-
speech_model=speech_model,
217-
language_code=language_code,
218-
speaker_labels=speaker_labels or None,
257+
transcription_config = _transcription_config(
258+
speech_model, language_code=language_code, speaker_labels=speaker_labels
219259
)
220260
results: list[_ItemResult] = []
221261
for index, item in enumerate(data.items, start=1):

‎tests/__snapshots__/test_cli_output_snapshots.ambr‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -277,8 +277,9 @@
277277
│ --text-column TEXT Reference text column │
278278
│ name (default: │
279279
│ auto-detect). │
280-
│ --speech-model [best|nano|slam-1|univer Speech model to │
281-
│ sal] evaluate. │
280+
│ --speech-model [universal-3-pro|univers Speech model to │
281+
│ al-2|slam-1|universal|be evaluate. │
282+
│ st|nano] │
282283
│ --language-code TEXT Force a language (e.g. │
283284
│ en_us). │
284285
│ --speaker-labels Diarize and also score │
@@ -302,7 +303,7 @@
302303
Score a model on 10 rows of an HF dataset
303304
$ assembly eval distil-whisper/meanwhile
304305
Compare models on your own audio
305-
$ assembly eval calls.csv --speech-model universal
306+
$ assembly eval calls.csv --speech-model universal-3-pro
306307
Score diarization too (WER + DER)
307308
$ assembly eval agent-calls.jsonl --speaker-labels
308309
Pick a subset/split and more rows

‎tests/test_eval_command.py‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -132,9 +132,33 @@ def test_speech_model_flag_reaches_config_and_output(tmp_path, mocker):
132132
assert _payload_of(result)["speech_model"] == "universal"
133133
tx_config = tx.call_args.kwargs["config"]
134134
assert tx_config.speech_model == aai.SpeechModel.universal
135+
assert tx_config.speech_models is None # legacy models ride the enum parameter only
135136
assert tx_config.speaker_labels is None # not requested -> omitted, not False
136137

137138

139+
@pytest.mark.parametrize("model", ["universal-3-pro", "universal-2"])
140+
def test_current_models_ride_the_speech_models_list(tmp_path, mocker, model):
141+
_auth()
142+
_write_wer_manifest(tmp_path)
143+
tx = _mock_transcribe(mocker, [_transcript("hello there"), _transcript("goodbye now")])
144+
result = runner.invoke(app, ["eval", "manifest.csv", "--speech-model", model, "--json"])
145+
assert result.exit_code == 0
146+
assert _payload_of(result)["speech_model"] == model
147+
tx_config = tx.call_args.kwargs["config"]
148+
assert tx_config.speech_models == [model]
149+
assert tx_config.speech_model is None # not in the SDK enum -> enum parameter omitted
150+
151+
152+
def test_no_speech_model_leaves_both_model_parameters_unset(tmp_path, mocker):
153+
_auth()
154+
_write_wer_manifest(tmp_path)
155+
tx = _mock_transcribe(mocker, [_transcript("hello there"), _transcript("goodbye now")])
156+
assert runner.invoke(app, ["eval", "manifest.csv"]).exit_code == 0
157+
tx_config = tx.call_args.kwargs["config"]
158+
assert tx_config.speech_model is None
159+
assert tx_config.speech_models is None
160+
161+
138162
def test_speech_model_named_in_human_header(tmp_path, mocker):
139163
_auth()
140164
_write_wer_manifest(tmp_path)

0 commit comments

Comments
 (0)