cli-modelarium 0.1.3__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.
- cli_modelarium/__init__.py +6 -0
- cli_modelarium/__main__.py +8 -0
- cli_modelarium/assertions.py +596 -0
- cli_modelarium/banner.py +96 -0
- cli_modelarium/batch.py +425 -0
- cli_modelarium/cli.py +2577 -0
- cli_modelarium/exceptions.py +88 -0
- cli_modelarium/hallucination.py +384 -0
- cli_modelarium/io_safety.py +112 -0
- cli_modelarium/judging.py +469 -0
- cli_modelarium/models_registry.py +138 -0
- cli_modelarium/output_formatters.py +1108 -0
- cli_modelarium/pricing.py +199 -0
- cli_modelarium/providers/__init__.py +7 -0
- cli_modelarium/providers/_utils.py +26 -0
- cli_modelarium/providers/anthropic_provider.py +148 -0
- cli_modelarium/providers/base.py +87 -0
- cli_modelarium/providers/deepseek_provider.py +15 -0
- cli_modelarium/providers/google_provider.py +135 -0
- cli_modelarium/providers/groq_provider.py +15 -0
- cli_modelarium/providers/local_provider.py +94 -0
- cli_modelarium/providers/mistral_provider.py +172 -0
- cli_modelarium/providers/openai_provider.py +163 -0
- cli_modelarium/providers/openrouter_provider.py +33 -0
- cli_modelarium/providers/xai_provider.py +15 -0
- cli_modelarium/run_statistics.py +1202 -0
- cli_modelarium/security.py +202 -0
- cli_modelarium/streaming.py +416 -0
- cli_modelarium-0.1.3.dist-info/METADATA +764 -0
- cli_modelarium-0.1.3.dist-info/RECORD +34 -0
- cli_modelarium-0.1.3.dist-info/WHEEL +4 -0
- cli_modelarium-0.1.3.dist-info/entry_points.txt +2 -0
- cli_modelarium-0.1.3.dist-info/licenses/LICENSE +201 -0
- cli_modelarium-0.1.3.dist-info/licenses/NOTICE +102 -0
cli_modelarium/cli.py
ADDED
|
@@ -0,0 +1,2577 @@
|
|
|
1
|
+
"""Click CLI entry point for Cli Modelarium."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import sys
|
|
7
|
+
from collections.abc import Callable
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
|
|
10
|
+
import click
|
|
11
|
+
import httpx
|
|
12
|
+
from rich.console import Console
|
|
13
|
+
from rich.panel import Panel
|
|
14
|
+
from rich.prompt import Prompt
|
|
15
|
+
from rich.table import Table
|
|
16
|
+
|
|
17
|
+
from cli_modelarium import __version__
|
|
18
|
+
from cli_modelarium.assertions import (
|
|
19
|
+
AssertionResult,
|
|
20
|
+
count_failed,
|
|
21
|
+
count_passed,
|
|
22
|
+
run_assertions,
|
|
23
|
+
)
|
|
24
|
+
from cli_modelarium.banner import render_banner, should_show_banner
|
|
25
|
+
from cli_modelarium.batch import (
|
|
26
|
+
ESTIMATE_INPUT_TOKENS,
|
|
27
|
+
ESTIMATE_OUTPUT_TOKENS,
|
|
28
|
+
BatchPrompt,
|
|
29
|
+
build_batch_states,
|
|
30
|
+
check_batch_size_limits,
|
|
31
|
+
detect_output_format,
|
|
32
|
+
estimate_batch_cost,
|
|
33
|
+
estimate_compare_cost,
|
|
34
|
+
load_batch_file,
|
|
35
|
+
output_overlaps_input,
|
|
36
|
+
run_batch,
|
|
37
|
+
)
|
|
38
|
+
from cli_modelarium.exceptions import (
|
|
39
|
+
BatchSizeError,
|
|
40
|
+
BatchValidationError,
|
|
41
|
+
KeyNotConfiguredError,
|
|
42
|
+
ModelariumError,
|
|
43
|
+
OutputFormatError,
|
|
44
|
+
UnknownModelError,
|
|
45
|
+
UnknownProviderError,
|
|
46
|
+
)
|
|
47
|
+
from cli_modelarium.hallucination import (
|
|
48
|
+
HALLUCINATION_TOS_EXTENSION,
|
|
49
|
+
annotate_risk_levels,
|
|
50
|
+
parse_hallucination_response,
|
|
51
|
+
resolve_hallucination_config,
|
|
52
|
+
)
|
|
53
|
+
from cli_modelarium.io_safety import load_system_prompt, split_escaped_csv
|
|
54
|
+
from cli_modelarium.judging import (
|
|
55
|
+
DEFAULT_CRITERIA,
|
|
56
|
+
JUDGE_PROMPT_TEMPLATE,
|
|
57
|
+
JudgeResult,
|
|
58
|
+
print_tos_disclosure,
|
|
59
|
+
run_judging,
|
|
60
|
+
total_judge_calls,
|
|
61
|
+
total_judge_cost,
|
|
62
|
+
)
|
|
63
|
+
from cli_modelarium.models_registry import (
|
|
64
|
+
all_known_providers,
|
|
65
|
+
list_models_for_provider,
|
|
66
|
+
parse_models_arg,
|
|
67
|
+
)
|
|
68
|
+
from cli_modelarium.output_formatters import (
|
|
69
|
+
BatchResult,
|
|
70
|
+
render_markdown_to_console,
|
|
71
|
+
state_to_result,
|
|
72
|
+
write_csv,
|
|
73
|
+
write_json,
|
|
74
|
+
write_markdown,
|
|
75
|
+
)
|
|
76
|
+
from cli_modelarium.pricing import (
|
|
77
|
+
PRICING,
|
|
78
|
+
is_local_model,
|
|
79
|
+
pricing_freshness_note,
|
|
80
|
+
)
|
|
81
|
+
from cli_modelarium.providers.base import BaseProvider
|
|
82
|
+
from cli_modelarium.providers.local_provider import LocalProvider
|
|
83
|
+
from cli_modelarium.security import (
|
|
84
|
+
KEY_PATTERNS,
|
|
85
|
+
delete_key,
|
|
86
|
+
delete_local_url,
|
|
87
|
+
is_key_configured,
|
|
88
|
+
load_local_url,
|
|
89
|
+
redact_secrets,
|
|
90
|
+
save_key,
|
|
91
|
+
save_local_url,
|
|
92
|
+
)
|
|
93
|
+
from cli_modelarium.streaming import (
|
|
94
|
+
DEFAULT_CONCURRENCY,
|
|
95
|
+
StreamState,
|
|
96
|
+
run_streaming_comparison,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
# Lazy provider import map. Each value is `module_path:ClassName` and is
|
|
100
|
+
# resolved via importlib at call time so we don't pay for every SDK import
|
|
101
|
+
# on every CLI invocation (matters for fast `--help` and `list-models`).
|
|
102
|
+
PROVIDER_REGISTRY: dict[str, str] = {
|
|
103
|
+
"openai": "cli_modelarium.providers.openai_provider:OpenAIProvider",
|
|
104
|
+
"anthropic": "cli_modelarium.providers.anthropic_provider:AnthropicProvider",
|
|
105
|
+
"google": "cli_modelarium.providers.google_provider:GoogleProvider",
|
|
106
|
+
"xai": "cli_modelarium.providers.xai_provider:XAIProvider",
|
|
107
|
+
"deepseek": "cli_modelarium.providers.deepseek_provider:DeepSeekProvider",
|
|
108
|
+
"groq": "cli_modelarium.providers.groq_provider:GroqProvider",
|
|
109
|
+
"openrouter": "cli_modelarium.providers.openrouter_provider:OpenRouterProvider",
|
|
110
|
+
"mistral": "cli_modelarium.providers.mistral_provider:MistralProvider",
|
|
111
|
+
"local": "cli_modelarium.providers.local_provider:LocalProvider",
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
# Exit codes used across the CLI (matches CI/CD conventions).
|
|
115
|
+
EXIT_OK = 0
|
|
116
|
+
EXIT_ASSERTION_FAILED = 1 # batch mode: at least one assertion failed
|
|
117
|
+
EXIT_CALL_FAILED = 2
|
|
118
|
+
|
|
119
|
+
console = Console()
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class _DefaultCommandGroup(click.Group):
|
|
123
|
+
"""Routes unknown bare arguments to the `compare` subcommand.
|
|
124
|
+
|
|
125
|
+
Lets users run `cli-modelarium "prompt" --models X` without typing the
|
|
126
|
+
`compare` verb explicitly.
|
|
127
|
+
"""
|
|
128
|
+
|
|
129
|
+
def resolve_command(
|
|
130
|
+
self, ctx: click.Context, args: list[str]
|
|
131
|
+
) -> tuple[str | None, click.Command | None, list[str]]:
|
|
132
|
+
try:
|
|
133
|
+
return super().resolve_command(ctx, args)
|
|
134
|
+
except click.UsageError:
|
|
135
|
+
if args and not args[0].startswith("-"):
|
|
136
|
+
args.insert(0, "compare")
|
|
137
|
+
return super().resolve_command(ctx, args)
|
|
138
|
+
raise
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
@click.group(
|
|
142
|
+
cls=_DefaultCommandGroup,
|
|
143
|
+
invoke_without_command=True,
|
|
144
|
+
context_settings={"help_option_names": ["-h", "--help"]},
|
|
145
|
+
)
|
|
146
|
+
@click.version_option(version=__version__, prog_name="cli-modelarium")
|
|
147
|
+
@click.pass_context
|
|
148
|
+
def main(ctx: click.Context) -> None:
|
|
149
|
+
"""Cli Modelarium - compare LLM outputs side-by-side from your terminal."""
|
|
150
|
+
if ctx.invoked_subcommand is None:
|
|
151
|
+
if should_show_banner():
|
|
152
|
+
render_banner()
|
|
153
|
+
click.echo(ctx.get_help())
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
# ===== compare =====
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
@main.command()
|
|
160
|
+
@click.argument("prompt", required=True)
|
|
161
|
+
@click.option("--models", required=True, help="Comma-separated model IDs or group names.")
|
|
162
|
+
@click.option("--temperatures", default="0.0", help="Comma-separated temperatures (default: 0.0).")
|
|
163
|
+
@click.option("--system-prompt", help="System prompt applied to every model.")
|
|
164
|
+
@click.option(
|
|
165
|
+
"--system-prompts",
|
|
166
|
+
help=(
|
|
167
|
+
"Comma-separated system prompts; the comparison fans out across them. "
|
|
168
|
+
"Use \\, for a literal comma."
|
|
169
|
+
),
|
|
170
|
+
)
|
|
171
|
+
@click.option(
|
|
172
|
+
"--system-prompt-file",
|
|
173
|
+
type=click.Path(),
|
|
174
|
+
help="Load a single system prompt from a UTF-8 file (max 1 MB).",
|
|
175
|
+
)
|
|
176
|
+
@click.option("--judge", help="Score outputs using this model as judge.")
|
|
177
|
+
@click.option(
|
|
178
|
+
"--judges",
|
|
179
|
+
help="Comma-separated panel of judges (scores averaged). Use \\, for a literal comma.",
|
|
180
|
+
)
|
|
181
|
+
@click.option(
|
|
182
|
+
"--judge-criteria",
|
|
183
|
+
help="Comma-separated custom scoring criteria. Use \\, for a literal comma.",
|
|
184
|
+
)
|
|
185
|
+
@click.option(
|
|
186
|
+
"--judge-template",
|
|
187
|
+
type=click.Path(),
|
|
188
|
+
help="Load a custom judge prompt template from a UTF-8 file (max 1 MB).",
|
|
189
|
+
)
|
|
190
|
+
@click.option(
|
|
191
|
+
"--include-reasoning", is_flag=True, help="Show each judge's reasoning in the output."
|
|
192
|
+
)
|
|
193
|
+
@click.option(
|
|
194
|
+
"--no-judge-tos",
|
|
195
|
+
is_flag=True,
|
|
196
|
+
help="Suppress the judge-use ToS reminder (for CI/CD where it's been acknowledged).",
|
|
197
|
+
)
|
|
198
|
+
@click.option(
|
|
199
|
+
"--check-hallucination",
|
|
200
|
+
is_flag=True,
|
|
201
|
+
help="Apply the hallucination detection preset. Requires --judge or --judges.",
|
|
202
|
+
)
|
|
203
|
+
@click.option(
|
|
204
|
+
"--expected-facts",
|
|
205
|
+
help="Comma-separated reference facts for hallucination check. Use \\, for a literal comma.",
|
|
206
|
+
)
|
|
207
|
+
@click.option(
|
|
208
|
+
"--expected-facts-file",
|
|
209
|
+
type=click.Path(),
|
|
210
|
+
help="Load expected facts from a .txt (one per line) or .json (array of strings) file.",
|
|
211
|
+
)
|
|
212
|
+
@click.option(
|
|
213
|
+
"--hallucination-template",
|
|
214
|
+
type=click.Path(),
|
|
215
|
+
help="Override the hallucination criteria text with a custom UTF-8 file (max 1 MB).",
|
|
216
|
+
)
|
|
217
|
+
@click.option(
|
|
218
|
+
"--output",
|
|
219
|
+
type=click.Path(),
|
|
220
|
+
help=(
|
|
221
|
+
"Write results to this file. Format inferred from extension "
|
|
222
|
+
"(.csv, .json, .md). Default: render Rich table to stdout."
|
|
223
|
+
),
|
|
224
|
+
)
|
|
225
|
+
@click.option(
|
|
226
|
+
"--output-format",
|
|
227
|
+
type=click.Choice(["csv", "json", "markdown"], case_sensitive=False),
|
|
228
|
+
help="Override the output format inferred from --output's extension (csv | json | markdown).",
|
|
229
|
+
)
|
|
230
|
+
@click.option(
|
|
231
|
+
"--max-cost",
|
|
232
|
+
type=click.FloatRange(min=0.0),
|
|
233
|
+
help="Refuse to run if estimated cost exceeds this USD (excludes judge cost).",
|
|
234
|
+
)
|
|
235
|
+
@click.option(
|
|
236
|
+
"--force",
|
|
237
|
+
is_flag=True,
|
|
238
|
+
help="Overwrite the output file if it exists.",
|
|
239
|
+
)
|
|
240
|
+
@click.option(
|
|
241
|
+
"--concurrency",
|
|
242
|
+
type=int,
|
|
243
|
+
default=DEFAULT_CONCURRENCY,
|
|
244
|
+
help=f"Max concurrent calls per provider (default: {DEFAULT_CONCURRENCY}).",
|
|
245
|
+
)
|
|
246
|
+
@click.option(
|
|
247
|
+
"--local-url",
|
|
248
|
+
help=(
|
|
249
|
+
"Override the default URL for the local model server "
|
|
250
|
+
"(default: http://localhost:11434/v1, Ollama)."
|
|
251
|
+
),
|
|
252
|
+
)
|
|
253
|
+
@click.option("--no-stream", is_flag=True, help="Disable live streaming display.")
|
|
254
|
+
@click.option(
|
|
255
|
+
"--runs",
|
|
256
|
+
type=click.IntRange(1, 100),
|
|
257
|
+
default=1,
|
|
258
|
+
show_default=True,
|
|
259
|
+
help=(
|
|
260
|
+
"Number of times to run each (model, temperature, system_prompt) "
|
|
261
|
+
"combination. Statistical analysis shown when > 1. Range: 1-100. "
|
|
262
|
+
"Cost multiplies by this value - use --max-cost for safety."
|
|
263
|
+
),
|
|
264
|
+
)
|
|
265
|
+
@click.option(
|
|
266
|
+
"--show-all-runs",
|
|
267
|
+
is_flag=True,
|
|
268
|
+
help=(
|
|
269
|
+
"Override the auto-collapse heuristic that disables the live display "
|
|
270
|
+
"when --runs creates more than 12 concurrent tasks. Forces every "
|
|
271
|
+
"run to render its own streaming panel."
|
|
272
|
+
),
|
|
273
|
+
)
|
|
274
|
+
@click.option(
|
|
275
|
+
"--significance/--no-significance",
|
|
276
|
+
default=None,
|
|
277
|
+
help=(
|
|
278
|
+
"Compute pairwise statistical significance tests between models. "
|
|
279
|
+
"Auto-enabled when --runs > 1 with 2+ models. Use --no-significance "
|
|
280
|
+
"to disable."
|
|
281
|
+
),
|
|
282
|
+
)
|
|
283
|
+
@click.option(
|
|
284
|
+
"--significance-threshold",
|
|
285
|
+
type=click.FloatRange(0.0, 1.0, min_open=True, max_open=True),
|
|
286
|
+
default=0.05,
|
|
287
|
+
show_default=True,
|
|
288
|
+
help=(
|
|
289
|
+
"P-value threshold for declaring significance. Common values: "
|
|
290
|
+
"0.05 (default), 0.01 (strict), 0.001 (very strict)."
|
|
291
|
+
),
|
|
292
|
+
)
|
|
293
|
+
@click.option(
|
|
294
|
+
"--significance-test",
|
|
295
|
+
type=click.Choice(["welch", "mann-whitney", "paired-t", "wilcoxon-signed"]),
|
|
296
|
+
default="welch",
|
|
297
|
+
show_default=True,
|
|
298
|
+
help=(
|
|
299
|
+
"Statistical test to use. 'welch' (default) handles unequal "
|
|
300
|
+
"variances. 'mann-whitney' is non-parametric (no normality "
|
|
301
|
+
"assumption). 'paired-t' uses scipy.stats.ttest_rel for "
|
|
302
|
+
"same-prompt paired comparisons (more statistical power). "
|
|
303
|
+
"'wilcoxon-signed' is the non-parametric paired alternative."
|
|
304
|
+
),
|
|
305
|
+
)
|
|
306
|
+
@click.option(
|
|
307
|
+
"--correction",
|
|
308
|
+
type=click.Choice(["none", "bonferroni", "holm"]),
|
|
309
|
+
default="bonferroni",
|
|
310
|
+
show_default=True,
|
|
311
|
+
help=(
|
|
312
|
+
"Multiple comparison correction. 'bonferroni' (default) is "
|
|
313
|
+
"conservative. 'holm' is less conservative while still "
|
|
314
|
+
"controlling family-wise error rate. 'none' is risky with 3+ "
|
|
315
|
+
"models."
|
|
316
|
+
),
|
|
317
|
+
)
|
|
318
|
+
@click.option(
|
|
319
|
+
"--significance-metric",
|
|
320
|
+
type=click.Choice(["score", "latency_ms", "output_tokens", "cost_usd"]),
|
|
321
|
+
default=None,
|
|
322
|
+
help=(
|
|
323
|
+
"Metric to test for significance. Default: 'score' when --judge "
|
|
324
|
+
"enabled, 'latency_ms' otherwise."
|
|
325
|
+
),
|
|
326
|
+
)
|
|
327
|
+
@click.option(
|
|
328
|
+
"--confidence-intervals/--no-confidence-intervals",
|
|
329
|
+
default=None,
|
|
330
|
+
help=(
|
|
331
|
+
"Compute bootstrap confidence intervals on per-cell means. "
|
|
332
|
+
"Auto-enabled when --runs > 1. Use --no-confidence-intervals to "
|
|
333
|
+
"disable."
|
|
334
|
+
),
|
|
335
|
+
)
|
|
336
|
+
@click.option(
|
|
337
|
+
"--ci-level",
|
|
338
|
+
type=click.FloatRange(0.0, 1.0, min_open=True, max_open=True),
|
|
339
|
+
default=0.95,
|
|
340
|
+
show_default=True,
|
|
341
|
+
help="Confidence level for bootstrap CIs (e.g. 0.95 for 95% CI).",
|
|
342
|
+
)
|
|
343
|
+
@click.option(
|
|
344
|
+
"--ci-method",
|
|
345
|
+
type=click.Choice(["bca", "percentile", "basic"]),
|
|
346
|
+
default="bca",
|
|
347
|
+
show_default=True,
|
|
348
|
+
help=(
|
|
349
|
+
"Bootstrap CI method. 'bca' (default) is bias-corrected and "
|
|
350
|
+
"accelerated - the publication-grade standard. 'percentile' is "
|
|
351
|
+
"simpler but less accurate near distribution tails. 'basic' is "
|
|
352
|
+
"the reverse-percentile method."
|
|
353
|
+
),
|
|
354
|
+
)
|
|
355
|
+
@click.option(
|
|
356
|
+
"--bootstrap-resamples",
|
|
357
|
+
type=click.IntRange(min=100),
|
|
358
|
+
default=5000,
|
|
359
|
+
show_default=True,
|
|
360
|
+
help=(
|
|
361
|
+
"Number of bootstrap resamples for CI computation. "
|
|
362
|
+
"Publication standard: 5000. Faster: 1000. More accurate: 10000."
|
|
363
|
+
),
|
|
364
|
+
)
|
|
365
|
+
@click.option(
|
|
366
|
+
"--bootstrap-seed",
|
|
367
|
+
type=int,
|
|
368
|
+
default=None,
|
|
369
|
+
help=(
|
|
370
|
+
"Random seed for reproducible bootstrap CIs. REQUIRED for "
|
|
371
|
+
"publication-grade output - without a seed, CIs vary slightly "
|
|
372
|
+
"across invocations."
|
|
373
|
+
),
|
|
374
|
+
)
|
|
375
|
+
def compare(
|
|
376
|
+
prompt: str,
|
|
377
|
+
models: str,
|
|
378
|
+
temperatures: str,
|
|
379
|
+
system_prompt: str | None,
|
|
380
|
+
system_prompts: str | None,
|
|
381
|
+
system_prompt_file: str | None,
|
|
382
|
+
judge: str | None,
|
|
383
|
+
judges: str | None,
|
|
384
|
+
judge_criteria: str | None,
|
|
385
|
+
judge_template: str | None,
|
|
386
|
+
include_reasoning: bool,
|
|
387
|
+
no_judge_tos: bool,
|
|
388
|
+
check_hallucination: bool,
|
|
389
|
+
expected_facts: str | None,
|
|
390
|
+
expected_facts_file: str | None,
|
|
391
|
+
hallucination_template: str | None,
|
|
392
|
+
output: str | None,
|
|
393
|
+
output_format: str | None,
|
|
394
|
+
max_cost: float | None,
|
|
395
|
+
force: bool,
|
|
396
|
+
concurrency: int,
|
|
397
|
+
local_url: str | None,
|
|
398
|
+
no_stream: bool,
|
|
399
|
+
runs: int,
|
|
400
|
+
show_all_runs: bool,
|
|
401
|
+
significance: bool | None,
|
|
402
|
+
significance_threshold: float,
|
|
403
|
+
significance_test: str,
|
|
404
|
+
correction: str,
|
|
405
|
+
significance_metric: str | None,
|
|
406
|
+
confidence_intervals: bool | None,
|
|
407
|
+
ci_level: float,
|
|
408
|
+
ci_method: str,
|
|
409
|
+
bootstrap_resamples: int,
|
|
410
|
+
bootstrap_seed: int | None,
|
|
411
|
+
) -> None:
|
|
412
|
+
"""Run a side-by-side comparison of LLMs on a single prompt."""
|
|
413
|
+
try:
|
|
414
|
+
model_list = parse_models_arg(models)
|
|
415
|
+
if not model_list:
|
|
416
|
+
raise click.UsageError("--models must include at least one model ID or group.")
|
|
417
|
+
temp_list = _parse_temperatures(temperatures)
|
|
418
|
+
system_prompt_list = _resolve_system_prompts(
|
|
419
|
+
system_prompt=system_prompt,
|
|
420
|
+
system_prompts=system_prompts,
|
|
421
|
+
system_prompt_file=system_prompt_file,
|
|
422
|
+
)
|
|
423
|
+
judge_models = _resolve_judge_models(judge=judge, judges=judges)
|
|
424
|
+
judge_criteria_list, judge_template_text = _resolve_judge_criteria_and_template(
|
|
425
|
+
judge_criteria=judge_criteria,
|
|
426
|
+
judge_template=judge_template,
|
|
427
|
+
)
|
|
428
|
+
# Phase 10 hallucination preset overrides judge criteria + template
|
|
429
|
+
# AND swaps in the hallucination response parser. None when not active.
|
|
430
|
+
hallucination_config = resolve_hallucination_config(
|
|
431
|
+
check_hallucination=check_hallucination,
|
|
432
|
+
expected_facts=expected_facts,
|
|
433
|
+
expected_facts_file=expected_facts_file,
|
|
434
|
+
hallucination_template=hallucination_template,
|
|
435
|
+
judge_models_present=bool(judge_models),
|
|
436
|
+
)
|
|
437
|
+
if hallucination_config is not None:
|
|
438
|
+
judge_criteria_list = hallucination_config.criteria
|
|
439
|
+
judge_template_text = hallucination_config.template
|
|
440
|
+
# Validate judge models BEFORE the main comparison runs - misconfigured
|
|
441
|
+
# judges should not fail late after burning money on the comparison.
|
|
442
|
+
if judge_models:
|
|
443
|
+
_validate_judge_models(judge_models, local_url=local_url)
|
|
444
|
+
except click.UsageError:
|
|
445
|
+
raise
|
|
446
|
+
except (
|
|
447
|
+
UnknownModelError,
|
|
448
|
+
KeyNotConfiguredError,
|
|
449
|
+
BatchValidationError,
|
|
450
|
+
FileNotFoundError,
|
|
451
|
+
ValueError,
|
|
452
|
+
) as e:
|
|
453
|
+
_print_error(str(e))
|
|
454
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
455
|
+
|
|
456
|
+
# Resolve --output / --output-format up front so a misconfigured path
|
|
457
|
+
# fails before we burn any API calls. output_path is None when the
|
|
458
|
+
# caller wants the default Rich display.
|
|
459
|
+
try:
|
|
460
|
+
output_path, output_fmt = _resolve_output_path(output, output_format, force)
|
|
461
|
+
except OutputFormatError as e:
|
|
462
|
+
_print_error(str(e))
|
|
463
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
464
|
+
|
|
465
|
+
# --max-cost pre-flight (excludes judge cost, matching batch).
|
|
466
|
+
# With --runs N, multiply the estimate by N before checking the ceiling.
|
|
467
|
+
if max_cost is not None:
|
|
468
|
+
per_run = estimate_compare_cost(model_list, temp_list, system_prompt_list)
|
|
469
|
+
estimated_total = per_run * runs
|
|
470
|
+
if estimated_total > max_cost:
|
|
471
|
+
if runs > 1:
|
|
472
|
+
_print_error(
|
|
473
|
+
f"Estimated cost ${estimated_total:.4f} "
|
|
474
|
+
f"(= ${per_run:.4f} x {runs} runs) exceeds --max-cost "
|
|
475
|
+
f"${max_cost:.4f}. Refusing to run."
|
|
476
|
+
)
|
|
477
|
+
else:
|
|
478
|
+
_print_error(
|
|
479
|
+
f"Estimated cost ${estimated_total:.4f} exceeds --max-cost "
|
|
480
|
+
f"${max_cost:.4f}. Refusing to run."
|
|
481
|
+
)
|
|
482
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
483
|
+
|
|
484
|
+
# Print a prominent cost warning when --runs > 1 is used without
|
|
485
|
+
# --max-cost, so the user is reminded that costs multiply by N.
|
|
486
|
+
if runs > 1 and max_cost is None:
|
|
487
|
+
per_run = estimate_compare_cost(model_list, temp_list, system_prompt_list)
|
|
488
|
+
estimated_total = per_run * runs
|
|
489
|
+
if estimated_total > 0:
|
|
490
|
+
console.print(
|
|
491
|
+
f"[yellow]Note: --runs {runs} multiplies cost. "
|
|
492
|
+
f"Estimated total: ${estimated_total:.4f} "
|
|
493
|
+
f"(= ${per_run:.4f} x {runs}).[/yellow]"
|
|
494
|
+
)
|
|
495
|
+
|
|
496
|
+
if judge_models and not no_judge_tos:
|
|
497
|
+
print_tos_disclosure(console)
|
|
498
|
+
if hallucination_config is not None:
|
|
499
|
+
console.print(
|
|
500
|
+
Panel(
|
|
501
|
+
HALLUCINATION_TOS_EXTENSION,
|
|
502
|
+
title="Hallucination detection",
|
|
503
|
+
border_style="yellow",
|
|
504
|
+
)
|
|
505
|
+
)
|
|
506
|
+
|
|
507
|
+
def provider_factory(name: str) -> BaseProvider:
|
|
508
|
+
return _get_provider_instance(name, local_url=local_url)
|
|
509
|
+
|
|
510
|
+
async def _run_all() -> tuple[list[StreamState], list[JudgeResult] | None]:
|
|
511
|
+
states = await run_streaming_comparison(
|
|
512
|
+
prompt=prompt,
|
|
513
|
+
models=model_list,
|
|
514
|
+
temperatures=temp_list,
|
|
515
|
+
system_prompts=system_prompt_list,
|
|
516
|
+
provider_factory=provider_factory,
|
|
517
|
+
console=console,
|
|
518
|
+
concurrency=concurrency,
|
|
519
|
+
live_display=not no_stream,
|
|
520
|
+
runs=runs,
|
|
521
|
+
show_all_runs=show_all_runs,
|
|
522
|
+
)
|
|
523
|
+
jrs: list[JudgeResult] | None = None
|
|
524
|
+
if judge_models:
|
|
525
|
+
# Judging strategy with --runs N:
|
|
526
|
+
# * Default: mode-only - judge one canonical output per cell
|
|
527
|
+
# (cheap; answers "what does this model usually say?").
|
|
528
|
+
# * --check-hallucination: per-run - judge every run so we can
|
|
529
|
+
# compute the hallucination rate across N runs.
|
|
530
|
+
# * runs == 1: existing behavior, one judge call per state.
|
|
531
|
+
if runs > 1 and hallucination_config is None:
|
|
532
|
+
jrs = await _run_mode_only_judging(
|
|
533
|
+
states=states,
|
|
534
|
+
prompt=prompt,
|
|
535
|
+
judge_models=judge_models,
|
|
536
|
+
criteria=judge_criteria_list,
|
|
537
|
+
template=judge_template_text,
|
|
538
|
+
provider_factory=provider_factory,
|
|
539
|
+
concurrency=concurrency,
|
|
540
|
+
)
|
|
541
|
+
else:
|
|
542
|
+
jrs = await run_judging(
|
|
543
|
+
items=[(s, prompt) for s in states],
|
|
544
|
+
judge_models=judge_models,
|
|
545
|
+
criteria=judge_criteria_list,
|
|
546
|
+
provider_factory=provider_factory,
|
|
547
|
+
template=judge_template_text,
|
|
548
|
+
response_parser=(
|
|
549
|
+
parse_hallucination_response
|
|
550
|
+
if hallucination_config is not None
|
|
551
|
+
else None
|
|
552
|
+
),
|
|
553
|
+
skip_self_eval=True,
|
|
554
|
+
concurrency=concurrency,
|
|
555
|
+
)
|
|
556
|
+
if hallucination_config is not None:
|
|
557
|
+
annotate_risk_levels(jrs)
|
|
558
|
+
return states, jrs
|
|
559
|
+
|
|
560
|
+
try:
|
|
561
|
+
# Single asyncio.run keeps all httpx client cleanup on one event loop,
|
|
562
|
+
# so we don't get "Event loop is closed" warnings on shutdown.
|
|
563
|
+
states, judge_results = asyncio.run(_run_all())
|
|
564
|
+
except KeyNotConfiguredError as e:
|
|
565
|
+
_print_error(str(e))
|
|
566
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
567
|
+
except ModelariumError as e:
|
|
568
|
+
_print_error(redact_secrets(str(e)))
|
|
569
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
570
|
+
|
|
571
|
+
# Pairwise significance: auto-enable when runs > 1 with 2+ models,
|
|
572
|
+
# unless the user explicitly opted out with --no-significance.
|
|
573
|
+
if significance is None:
|
|
574
|
+
should_compute_significance = runs > 1 and len(model_list) >= 2
|
|
575
|
+
else:
|
|
576
|
+
should_compute_significance = significance
|
|
577
|
+
|
|
578
|
+
# v0.1.3: bootstrap CIs auto-enable when runs > 1 (matching significance
|
|
579
|
+
# pattern). User can opt out with --no-confidence-intervals.
|
|
580
|
+
if confidence_intervals is None:
|
|
581
|
+
should_compute_ci = runs > 1
|
|
582
|
+
else:
|
|
583
|
+
should_compute_ci = confidence_intervals
|
|
584
|
+
|
|
585
|
+
significance_results = None
|
|
586
|
+
stats_by_cell_with_ci: dict | None = None
|
|
587
|
+
mcnemar_results = None
|
|
588
|
+
methodology: dict | None = None
|
|
589
|
+
|
|
590
|
+
if runs > 1 and len(model_list) >= 1:
|
|
591
|
+
# Always tag judge results with their state id so paired/score
|
|
592
|
+
# extractors can match them back.
|
|
593
|
+
if judge_results is not None:
|
|
594
|
+
for state, jr in zip(states, judge_results, strict=True):
|
|
595
|
+
jr._state_id = id(state) # type: ignore[attr-defined]
|
|
596
|
+
|
|
597
|
+
states_by_model: dict[str, list[StreamState]] = {}
|
|
598
|
+
for state in states:
|
|
599
|
+
states_by_model.setdefault(state.model, []).append(state)
|
|
600
|
+
|
|
601
|
+
if should_compute_significance and len(model_list) >= 2:
|
|
602
|
+
from cli_modelarium.run_statistics import (
|
|
603
|
+
compute_significance_with_ci,
|
|
604
|
+
)
|
|
605
|
+
|
|
606
|
+
sig_metric = significance_metric
|
|
607
|
+
if sig_metric is None:
|
|
608
|
+
sig_metric = "score" if judge_results is not None else "latency_ms"
|
|
609
|
+
|
|
610
|
+
try:
|
|
611
|
+
significance_results = compute_significance_with_ci(
|
|
612
|
+
states_by_model,
|
|
613
|
+
judge_results,
|
|
614
|
+
metric=sig_metric,
|
|
615
|
+
test=significance_test, # type: ignore[arg-type]
|
|
616
|
+
correction=correction, # type: ignore[arg-type]
|
|
617
|
+
threshold=significance_threshold,
|
|
618
|
+
compute_ci=should_compute_ci,
|
|
619
|
+
ci_level=ci_level,
|
|
620
|
+
ci_method=ci_method,
|
|
621
|
+
n_resamples=bootstrap_resamples,
|
|
622
|
+
seed=bootstrap_seed,
|
|
623
|
+
)
|
|
624
|
+
except ValueError as e:
|
|
625
|
+
console.print(f"[yellow]Significance test skipped: {e}[/yellow]")
|
|
626
|
+
significance_results = None
|
|
627
|
+
|
|
628
|
+
if should_compute_ci:
|
|
629
|
+
from cli_modelarium.run_statistics import compute_stats_with_cis
|
|
630
|
+
|
|
631
|
+
cis = compute_stats_with_cis(
|
|
632
|
+
states_by_model,
|
|
633
|
+
judge_results,
|
|
634
|
+
ci_level=ci_level,
|
|
635
|
+
ci_method=ci_method,
|
|
636
|
+
n_resamples=bootstrap_resamples,
|
|
637
|
+
seed=bootstrap_seed,
|
|
638
|
+
)
|
|
639
|
+
stats_by_cell_with_ci = _flatten_cell_cis(cis)
|
|
640
|
+
|
|
641
|
+
if (
|
|
642
|
+
hallucination_config is not None
|
|
643
|
+
and len(model_list) >= 2
|
|
644
|
+
and judge_results is not None
|
|
645
|
+
):
|
|
646
|
+
from cli_modelarium.run_statistics import compute_mcnemar_pairwise
|
|
647
|
+
|
|
648
|
+
judge_by_state_id = {
|
|
649
|
+
id(state): jr
|
|
650
|
+
for state, jr in zip(states, judge_results, strict=True)
|
|
651
|
+
}
|
|
652
|
+
mcnemar_results = compute_mcnemar_pairwise(
|
|
653
|
+
states_by_model,
|
|
654
|
+
judge_by_state_id,
|
|
655
|
+
correction=correction, # type: ignore[arg-type]
|
|
656
|
+
threshold=significance_threshold,
|
|
657
|
+
)
|
|
658
|
+
|
|
659
|
+
# Record methodology metadata for reproducibility.
|
|
660
|
+
import scipy as _scipy
|
|
661
|
+
|
|
662
|
+
methodology = {
|
|
663
|
+
"tool_version": __version__,
|
|
664
|
+
"scipy_version": _scipy.__version__,
|
|
665
|
+
"python_version": f"{sys.version_info.major}.{sys.version_info.minor}",
|
|
666
|
+
"n_runs": runs,
|
|
667
|
+
"bootstrap": {
|
|
668
|
+
"enabled": should_compute_ci,
|
|
669
|
+
"method": ci_method if should_compute_ci else None,
|
|
670
|
+
"n_resamples": bootstrap_resamples if should_compute_ci else None,
|
|
671
|
+
"ci_level": ci_level if should_compute_ci else None,
|
|
672
|
+
"seed": bootstrap_seed if should_compute_ci else None,
|
|
673
|
+
},
|
|
674
|
+
"significance": {
|
|
675
|
+
"enabled": bool(significance_results),
|
|
676
|
+
"test": significance_test if significance_results else None,
|
|
677
|
+
"correction": correction if significance_results else None,
|
|
678
|
+
"threshold": significance_threshold if significance_results else None,
|
|
679
|
+
},
|
|
680
|
+
}
|
|
681
|
+
|
|
682
|
+
if output_path is not None or output_fmt is not None:
|
|
683
|
+
# File or explicit-format output path: serialize via batch's writers.
|
|
684
|
+
results = _states_to_compare_results(states, prompt, judge_results)
|
|
685
|
+
_emit_batch_results(
|
|
686
|
+
results,
|
|
687
|
+
output_path=output_path,
|
|
688
|
+
output_fmt=output_fmt or "markdown",
|
|
689
|
+
runs=runs,
|
|
690
|
+
significance_results=significance_results,
|
|
691
|
+
stats_by_cell_cis=stats_by_cell_with_ci,
|
|
692
|
+
mcnemar_results=mcnemar_results,
|
|
693
|
+
methodology=methodology,
|
|
694
|
+
)
|
|
695
|
+
elif runs > 1:
|
|
696
|
+
_display_results_with_runs(
|
|
697
|
+
states,
|
|
698
|
+
judge_results=judge_results,
|
|
699
|
+
runs=runs,
|
|
700
|
+
include_reasoning=include_reasoning,
|
|
701
|
+
hallucination_mode=hallucination_config is not None,
|
|
702
|
+
hallucination_facts=(hallucination_config.facts if hallucination_config else None),
|
|
703
|
+
significance_results=significance_results,
|
|
704
|
+
stats_by_cell_cis=stats_by_cell_with_ci,
|
|
705
|
+
mcnemar_results=mcnemar_results,
|
|
706
|
+
)
|
|
707
|
+
else:
|
|
708
|
+
_display_results(
|
|
709
|
+
states,
|
|
710
|
+
judge_results=judge_results,
|
|
711
|
+
include_reasoning=include_reasoning,
|
|
712
|
+
hallucination_mode=hallucination_config is not None,
|
|
713
|
+
hallucination_facts=(hallucination_config.facts if hallucination_config else None),
|
|
714
|
+
)
|
|
715
|
+
|
|
716
|
+
if any(s.error for s in states):
|
|
717
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
718
|
+
sys.exit(EXIT_OK)
|
|
719
|
+
|
|
720
|
+
|
|
721
|
+
# ===== batch =====
|
|
722
|
+
|
|
723
|
+
|
|
724
|
+
@main.command()
|
|
725
|
+
@click.argument("file", type=click.Path(exists=True, dir_okay=False))
|
|
726
|
+
@click.option("--models", required=True, help="Comma-separated model IDs or group names.")
|
|
727
|
+
@click.option("--temperatures", default="0.0", help="Comma-separated temperatures (default: 0.0).")
|
|
728
|
+
@click.option(
|
|
729
|
+
"--system-prompt",
|
|
730
|
+
help=(
|
|
731
|
+
"System prompt applied to every prompt (per-prompt 'system' field in "
|
|
732
|
+
"the input file wins for that prompt)."
|
|
733
|
+
),
|
|
734
|
+
)
|
|
735
|
+
@click.option(
|
|
736
|
+
"--system-prompts",
|
|
737
|
+
help=(
|
|
738
|
+
"Comma-separated system prompts; the matrix fans out across them. "
|
|
739
|
+
"Use \\, for a literal comma."
|
|
740
|
+
),
|
|
741
|
+
)
|
|
742
|
+
@click.option(
|
|
743
|
+
"--system-prompt-file",
|
|
744
|
+
type=click.Path(),
|
|
745
|
+
help="Load a single system prompt from a UTF-8 file (max 1 MB).",
|
|
746
|
+
)
|
|
747
|
+
@click.option("--judge", help="Score outputs using this model as judge.")
|
|
748
|
+
@click.option(
|
|
749
|
+
"--judges",
|
|
750
|
+
help="Comma-separated panel of judges (scores averaged). Use \\, for a literal comma.",
|
|
751
|
+
)
|
|
752
|
+
@click.option(
|
|
753
|
+
"--judge-criteria",
|
|
754
|
+
help="Comma-separated custom scoring criteria. Use \\, for a literal comma.",
|
|
755
|
+
)
|
|
756
|
+
@click.option(
|
|
757
|
+
"--judge-template",
|
|
758
|
+
type=click.Path(),
|
|
759
|
+
help="Load a custom judge prompt template from a UTF-8 file (max 1 MB).",
|
|
760
|
+
)
|
|
761
|
+
@click.option(
|
|
762
|
+
"--include-reasoning", is_flag=True, help="Show each judge's reasoning in the output."
|
|
763
|
+
)
|
|
764
|
+
@click.option(
|
|
765
|
+
"--no-judge-tos",
|
|
766
|
+
is_flag=True,
|
|
767
|
+
help="Suppress the judge-use ToS reminder (for CI/CD where it's been acknowledged).",
|
|
768
|
+
)
|
|
769
|
+
@click.option(
|
|
770
|
+
"--check-hallucination",
|
|
771
|
+
is_flag=True,
|
|
772
|
+
help="Apply the hallucination detection preset. Requires --judge or --judges.",
|
|
773
|
+
)
|
|
774
|
+
@click.option(
|
|
775
|
+
"--expected-facts",
|
|
776
|
+
help="Comma-separated reference facts. Use \\, for a literal comma.",
|
|
777
|
+
)
|
|
778
|
+
@click.option(
|
|
779
|
+
"--expected-facts-file",
|
|
780
|
+
type=click.Path(),
|
|
781
|
+
help="Load expected facts from a .txt (one per line) or .json (array of strings) file.",
|
|
782
|
+
)
|
|
783
|
+
@click.option(
|
|
784
|
+
"--hallucination-template",
|
|
785
|
+
type=click.Path(),
|
|
786
|
+
help="Override the hallucination criteria text with a custom UTF-8 file (max 1 MB).",
|
|
787
|
+
)
|
|
788
|
+
@click.option(
|
|
789
|
+
"--output",
|
|
790
|
+
type=click.Path(),
|
|
791
|
+
help=(
|
|
792
|
+
"Output file path. Format auto-detected from extension; omit to render Markdown on stdout."
|
|
793
|
+
),
|
|
794
|
+
)
|
|
795
|
+
@click.option(
|
|
796
|
+
"--output-format",
|
|
797
|
+
type=click.Choice(["csv", "json", "markdown"], case_sensitive=False),
|
|
798
|
+
help="Override the output format inferred from --output extension.",
|
|
799
|
+
)
|
|
800
|
+
@click.option(
|
|
801
|
+
"--max-cost",
|
|
802
|
+
type=click.FloatRange(min=0.0),
|
|
803
|
+
help="Refuse to run if the estimated cost exceeds this USD (excludes judge cost).",
|
|
804
|
+
)
|
|
805
|
+
@click.option(
|
|
806
|
+
"--concurrency",
|
|
807
|
+
type=int,
|
|
808
|
+
default=DEFAULT_CONCURRENCY,
|
|
809
|
+
help=f"Max concurrent calls per provider (default: {DEFAULT_CONCURRENCY}).",
|
|
810
|
+
)
|
|
811
|
+
@click.option("--local-url", help="Override the default URL for the local model server.")
|
|
812
|
+
@click.option(
|
|
813
|
+
"--min-pass-rate",
|
|
814
|
+
type=float,
|
|
815
|
+
help=(
|
|
816
|
+
"Exit 1 if assertion pass rate falls below this threshold (0.0-1.0). "
|
|
817
|
+
"Default behaviour without this flag is strict: ANY assertion failure "
|
|
818
|
+
"exits 1."
|
|
819
|
+
),
|
|
820
|
+
)
|
|
821
|
+
@click.option(
|
|
822
|
+
"--no-assertions",
|
|
823
|
+
is_flag=True,
|
|
824
|
+
help=(
|
|
825
|
+
"Skip assertion checks entirely. Pass/fail counts are zeroed and "
|
|
826
|
+
"exit code reflects only call status."
|
|
827
|
+
),
|
|
828
|
+
)
|
|
829
|
+
@click.option(
|
|
830
|
+
"--strict-assertions",
|
|
831
|
+
is_flag=True,
|
|
832
|
+
help=(
|
|
833
|
+
"Make the default strict behaviour explicit (any assertion failure "
|
|
834
|
+
"exits 1). Mutually exclusive with --min-pass-rate."
|
|
835
|
+
),
|
|
836
|
+
)
|
|
837
|
+
@click.option(
|
|
838
|
+
"--no-judge",
|
|
839
|
+
is_flag=True,
|
|
840
|
+
help="Skip judge scoring even if --judge or --judges is configured.",
|
|
841
|
+
)
|
|
842
|
+
@click.option("--force", is_flag=True, help="Overwrite the output file if it exists.")
|
|
843
|
+
@click.option(
|
|
844
|
+
"--force-large", is_flag=True, help="Bypass safety caps (max 1000 prompts, max 10000 calls)."
|
|
845
|
+
)
|
|
846
|
+
def batch(
|
|
847
|
+
file: str,
|
|
848
|
+
models: str,
|
|
849
|
+
temperatures: str,
|
|
850
|
+
system_prompt: str | None,
|
|
851
|
+
system_prompts: str | None,
|
|
852
|
+
system_prompt_file: str | None,
|
|
853
|
+
judge: str | None,
|
|
854
|
+
judges: str | None,
|
|
855
|
+
judge_criteria: str | None,
|
|
856
|
+
judge_template: str | None,
|
|
857
|
+
include_reasoning: bool,
|
|
858
|
+
no_judge_tos: bool,
|
|
859
|
+
check_hallucination: bool,
|
|
860
|
+
expected_facts: str | None,
|
|
861
|
+
expected_facts_file: str | None,
|
|
862
|
+
hallucination_template: str | None,
|
|
863
|
+
output: str | None,
|
|
864
|
+
output_format: str | None,
|
|
865
|
+
max_cost: float | None,
|
|
866
|
+
concurrency: int,
|
|
867
|
+
local_url: str | None,
|
|
868
|
+
min_pass_rate: float | None,
|
|
869
|
+
no_assertions: bool,
|
|
870
|
+
strict_assertions: bool,
|
|
871
|
+
no_judge: bool,
|
|
872
|
+
force: bool,
|
|
873
|
+
force_large: bool,
|
|
874
|
+
) -> None:
|
|
875
|
+
"""Run a multi-prompt batch evaluation from a file.
|
|
876
|
+
|
|
877
|
+
The input file is parsed by extension (.txt or .json). Output format is
|
|
878
|
+
inferred from --output's extension (.csv / .json / .md), or pass
|
|
879
|
+
--output-format to override. Omit --output to render Markdown to stdout.
|
|
880
|
+
|
|
881
|
+
Per-prompt system prompts: include `"system": "..."` in a JSON prompt
|
|
882
|
+
object to override the command-line system prompt for that one prompt.
|
|
883
|
+
|
|
884
|
+
Assertions (JSON input only): include `"assertions": [...]` on a prompt
|
|
885
|
+
object. Exit codes: 0 = all passed, 1 = assertion failure(s) or pass
|
|
886
|
+
rate below --min-pass-rate, 2 = call failure or IO error. Call failures
|
|
887
|
+
win over assertion failures (2 > 1).
|
|
888
|
+
"""
|
|
889
|
+
# --strict-assertions and --min-pass-rate are alternatives; combining
|
|
890
|
+
# them is ambiguous, so reject upfront.
|
|
891
|
+
if strict_assertions and min_pass_rate is not None:
|
|
892
|
+
raise click.UsageError("--strict-assertions and --min-pass-rate are mutually exclusive.")
|
|
893
|
+
if min_pass_rate is not None and not (0.0 <= min_pass_rate <= 1.0):
|
|
894
|
+
raise click.UsageError(
|
|
895
|
+
f"--min-pass-rate must be between 0.0 and 1.0 (got {min_pass_rate})."
|
|
896
|
+
)
|
|
897
|
+
|
|
898
|
+
try:
|
|
899
|
+
prompts = load_batch_file(file)
|
|
900
|
+
if not prompts:
|
|
901
|
+
console.print(
|
|
902
|
+
Panel(
|
|
903
|
+
f"No prompts to run - the file at {file} parsed as empty.",
|
|
904
|
+
title="Batch",
|
|
905
|
+
border_style="yellow",
|
|
906
|
+
)
|
|
907
|
+
)
|
|
908
|
+
return
|
|
909
|
+
|
|
910
|
+
model_list = parse_models_arg(models)
|
|
911
|
+
if not model_list:
|
|
912
|
+
raise click.UsageError("--models must include at least one model ID or group.")
|
|
913
|
+
temp_list = _parse_temperatures(temperatures)
|
|
914
|
+
command_sp_list = _resolve_system_prompts(
|
|
915
|
+
system_prompt=system_prompt,
|
|
916
|
+
system_prompts=system_prompts,
|
|
917
|
+
system_prompt_file=system_prompt_file,
|
|
918
|
+
)
|
|
919
|
+
judge_models = _resolve_judge_models(judge=judge, judges=judges)
|
|
920
|
+
judge_criteria_list, judge_template_text = _resolve_judge_criteria_and_template(
|
|
921
|
+
judge_criteria=judge_criteria,
|
|
922
|
+
judge_template=judge_template,
|
|
923
|
+
)
|
|
924
|
+
# Phase 10 hallucination preset; None when not active. Overrides
|
|
925
|
+
# judge criteria and template, and swaps in the hallucination
|
|
926
|
+
# response parser later.
|
|
927
|
+
hallucination_config = resolve_hallucination_config(
|
|
928
|
+
check_hallucination=check_hallucination,
|
|
929
|
+
expected_facts=expected_facts,
|
|
930
|
+
expected_facts_file=expected_facts_file,
|
|
931
|
+
hallucination_template=hallucination_template,
|
|
932
|
+
judge_models_present=bool(judge_models and not no_judge),
|
|
933
|
+
)
|
|
934
|
+
if hallucination_config is not None:
|
|
935
|
+
judge_criteria_list = hallucination_config.criteria
|
|
936
|
+
judge_template_text = hallucination_config.template
|
|
937
|
+
# Validate judge models BEFORE the batch starts - misconfigured
|
|
938
|
+
# judges should not fail late after burning batch money.
|
|
939
|
+
if judge_models and not no_judge:
|
|
940
|
+
_validate_judge_models(judge_models, local_url=local_url)
|
|
941
|
+
except click.UsageError:
|
|
942
|
+
raise
|
|
943
|
+
except (
|
|
944
|
+
BatchValidationError,
|
|
945
|
+
UnknownModelError,
|
|
946
|
+
KeyNotConfiguredError,
|
|
947
|
+
FileNotFoundError,
|
|
948
|
+
ValueError,
|
|
949
|
+
) as e:
|
|
950
|
+
_print_error(str(e))
|
|
951
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
952
|
+
|
|
953
|
+
# Resolve output target + format.
|
|
954
|
+
try:
|
|
955
|
+
output_path, output_fmt = _resolve_batch_output(
|
|
956
|
+
input_path=file,
|
|
957
|
+
output=output,
|
|
958
|
+
output_format=output_format,
|
|
959
|
+
force=force,
|
|
960
|
+
)
|
|
961
|
+
except OutputFormatError as e:
|
|
962
|
+
_print_error(str(e))
|
|
963
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
964
|
+
|
|
965
|
+
# Size limits.
|
|
966
|
+
try:
|
|
967
|
+
total = check_batch_size_limits(
|
|
968
|
+
prompts,
|
|
969
|
+
model_list,
|
|
970
|
+
temp_list,
|
|
971
|
+
command_sp_list,
|
|
972
|
+
force_large=force_large,
|
|
973
|
+
)
|
|
974
|
+
except BatchSizeError as e:
|
|
975
|
+
_print_error(str(e))
|
|
976
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
977
|
+
|
|
978
|
+
# Cost ceiling.
|
|
979
|
+
if max_cost is not None:
|
|
980
|
+
est = estimate_batch_cost(prompts, model_list, temp_list, command_sp_list)
|
|
981
|
+
if est > max_cost:
|
|
982
|
+
_print_error(
|
|
983
|
+
f"Estimated cost ${est:.4f} exceeds --max-cost ${max_cost:.4f}.\n"
|
|
984
|
+
f" Estimate assumes {ESTIMATE_INPUT_TOKENS} input + "
|
|
985
|
+
f"{ESTIMATE_OUTPUT_TOKENS} output tokens per call across "
|
|
986
|
+
f"{total} call{'s' if total != 1 else ''}.\n"
|
|
987
|
+
f" Reduce dimensions or raise --max-cost to proceed."
|
|
988
|
+
)
|
|
989
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
990
|
+
|
|
991
|
+
# Build states + run.
|
|
992
|
+
def provider_factory(name: str) -> BaseProvider:
|
|
993
|
+
return _get_provider_instance(name, local_url=local_url)
|
|
994
|
+
|
|
995
|
+
pairs = build_batch_states(prompts, model_list, temp_list, command_sp_list)
|
|
996
|
+
console.print(
|
|
997
|
+
f"[dim]Running batch: {total} call{'s' if total != 1 else ''} "
|
|
998
|
+
f"({len(prompts)} prompt{'s' if len(prompts) != 1 else ''} x "
|
|
999
|
+
f"{len(model_list)} model{'s' if len(model_list) != 1 else ''} x "
|
|
1000
|
+
f"{len(temp_list)} temperature{'s' if len(temp_list) != 1 else ''})[/dim]"
|
|
1001
|
+
)
|
|
1002
|
+
|
|
1003
|
+
if judge_models and not no_judge and not no_judge_tos:
|
|
1004
|
+
print_tos_disclosure(console)
|
|
1005
|
+
if hallucination_config is not None:
|
|
1006
|
+
console.print(
|
|
1007
|
+
Panel(
|
|
1008
|
+
HALLUCINATION_TOS_EXTENSION,
|
|
1009
|
+
title="Hallucination detection",
|
|
1010
|
+
border_style="yellow",
|
|
1011
|
+
)
|
|
1012
|
+
)
|
|
1013
|
+
|
|
1014
|
+
async def _run_batch_and_judge() -> list[JudgeResult] | None:
|
|
1015
|
+
await run_batch(
|
|
1016
|
+
pairs=pairs,
|
|
1017
|
+
provider_factory=provider_factory,
|
|
1018
|
+
console=console,
|
|
1019
|
+
concurrency=concurrency,
|
|
1020
|
+
show_progress=True,
|
|
1021
|
+
)
|
|
1022
|
+
if judge_models and not no_judge:
|
|
1023
|
+
jrs = await run_judging(
|
|
1024
|
+
items=[(s, bp.prompt) for s, bp in pairs],
|
|
1025
|
+
judge_models=judge_models,
|
|
1026
|
+
criteria=judge_criteria_list,
|
|
1027
|
+
provider_factory=provider_factory,
|
|
1028
|
+
template=judge_template_text,
|
|
1029
|
+
response_parser=(
|
|
1030
|
+
parse_hallucination_response if hallucination_config is not None else None
|
|
1031
|
+
),
|
|
1032
|
+
skip_self_eval=True,
|
|
1033
|
+
concurrency=concurrency,
|
|
1034
|
+
)
|
|
1035
|
+
if hallucination_config is not None:
|
|
1036
|
+
annotate_risk_levels(jrs)
|
|
1037
|
+
return jrs
|
|
1038
|
+
return None
|
|
1039
|
+
|
|
1040
|
+
try:
|
|
1041
|
+
# Single asyncio.run keeps all httpx client cleanup on one event loop.
|
|
1042
|
+
judge_results = asyncio.run(_run_batch_and_judge())
|
|
1043
|
+
except KeyNotConfiguredError as e:
|
|
1044
|
+
_print_error(str(e))
|
|
1045
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1046
|
+
except ModelariumError as e:
|
|
1047
|
+
_print_error(redact_secrets(str(e)))
|
|
1048
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1049
|
+
|
|
1050
|
+
# Run assertions per-state. Failed-call states skip assertion execution
|
|
1051
|
+
# (no real output to check). Successful-call states with no configured
|
|
1052
|
+
# assertions get an empty list (not None) so they show as "0/0" - which
|
|
1053
|
+
# is vacuously fine.
|
|
1054
|
+
assertion_results_per_state: list[list[AssertionResult] | None] = []
|
|
1055
|
+
for state, bp in pairs:
|
|
1056
|
+
if no_assertions or state.error:
|
|
1057
|
+
assertion_results_per_state.append(None)
|
|
1058
|
+
elif not bp.assertions:
|
|
1059
|
+
assertion_results_per_state.append([])
|
|
1060
|
+
else:
|
|
1061
|
+
assertion_results_per_state.append(
|
|
1062
|
+
run_assertions(
|
|
1063
|
+
output=state.text,
|
|
1064
|
+
latency_ms=state.latency_ms,
|
|
1065
|
+
cost_usd=state.cost_usd,
|
|
1066
|
+
assertions=bp.assertions,
|
|
1067
|
+
)
|
|
1068
|
+
)
|
|
1069
|
+
|
|
1070
|
+
# Convert StreamStates to BatchResults and emit.
|
|
1071
|
+
results = []
|
|
1072
|
+
for i, (state, bp) in enumerate(pairs):
|
|
1073
|
+
jr = judge_results[i] if judge_results is not None else None
|
|
1074
|
+
ar = assertion_results_per_state[i]
|
|
1075
|
+
results.append(state_to_result(state, bp, judge_result=jr, assertion_results=ar))
|
|
1076
|
+
_emit_batch_results(results, output_path=output_path, output_fmt=output_fmt)
|
|
1077
|
+
|
|
1078
|
+
failed = sum(1 for r in results if r.error)
|
|
1079
|
+
success = len(results) - failed
|
|
1080
|
+
total_cost = sum(r.cost_usd for r in results if r.error is None)
|
|
1081
|
+
|
|
1082
|
+
# Tally assertion outcomes. count_passed excludes `error` rows from
|
|
1083
|
+
# both numerator and denominator, so a missing-jsonschema doesn't
|
|
1084
|
+
# poison the pass rate or trigger exit 1.
|
|
1085
|
+
total_assertion_passed = 0
|
|
1086
|
+
total_assertion_definitive = 0
|
|
1087
|
+
total_assertion_failed = 0
|
|
1088
|
+
for ar in assertion_results_per_state:
|
|
1089
|
+
if ar is None:
|
|
1090
|
+
continue
|
|
1091
|
+
p, d = count_passed(ar)
|
|
1092
|
+
total_assertion_passed += p
|
|
1093
|
+
total_assertion_definitive += d
|
|
1094
|
+
total_assertion_failed += count_failed(ar)
|
|
1095
|
+
|
|
1096
|
+
summary_parts = [
|
|
1097
|
+
f"[green]{success} succeeded[/green]",
|
|
1098
|
+
f"[red]{failed} failed[/red]",
|
|
1099
|
+
f"[dim]total cost ${total_cost:.6f}[/dim]",
|
|
1100
|
+
]
|
|
1101
|
+
if judge_results is not None:
|
|
1102
|
+
j_cost = total_judge_cost(judge_results)
|
|
1103
|
+
j_calls = total_judge_calls(judge_results)
|
|
1104
|
+
summary_parts.append(
|
|
1105
|
+
f"[dim]judge cost ${j_cost:.6f} ({j_calls} call{'s' if j_calls != 1 else ''})[/dim]"
|
|
1106
|
+
)
|
|
1107
|
+
if total_assertion_definitive > 0 or total_assertion_failed > 0:
|
|
1108
|
+
pass_rate = (
|
|
1109
|
+
total_assertion_passed / total_assertion_definitive
|
|
1110
|
+
if total_assertion_definitive > 0
|
|
1111
|
+
else 1.0
|
|
1112
|
+
)
|
|
1113
|
+
rate_color = "green" if total_assertion_failed == 0 else "red"
|
|
1114
|
+
summary_parts.append(
|
|
1115
|
+
f"[{rate_color}]assertions {total_assertion_passed}/{total_assertion_definitive} "
|
|
1116
|
+
f"({pass_rate * 100:.0f}%)[/{rate_color}]"
|
|
1117
|
+
)
|
|
1118
|
+
console.print(" ".join(summary_parts))
|
|
1119
|
+
|
|
1120
|
+
# Exit-code logic. Call failures dominate - usually they mean the user
|
|
1121
|
+
# needs to fix credentials/infra before they can even evaluate
|
|
1122
|
+
# assertions, so we surface that as 2 rather than the softer 1.
|
|
1123
|
+
if failed > 0:
|
|
1124
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1125
|
+
|
|
1126
|
+
if not no_assertions:
|
|
1127
|
+
if min_pass_rate is not None:
|
|
1128
|
+
# --min-pass-rate threshold mode: tolerate some failures.
|
|
1129
|
+
if total_assertion_definitive > 0:
|
|
1130
|
+
pass_rate = total_assertion_passed / total_assertion_definitive
|
|
1131
|
+
if pass_rate < min_pass_rate:
|
|
1132
|
+
sys.exit(EXIT_ASSERTION_FAILED)
|
|
1133
|
+
else:
|
|
1134
|
+
# Default / --strict-assertions: ANY failure exits 1.
|
|
1135
|
+
if total_assertion_failed > 0:
|
|
1136
|
+
sys.exit(EXIT_ASSERTION_FAILED)
|
|
1137
|
+
|
|
1138
|
+
sys.exit(EXIT_OK)
|
|
1139
|
+
|
|
1140
|
+
|
|
1141
|
+
def _resolve_output_path(
|
|
1142
|
+
output: str | None,
|
|
1143
|
+
output_format: str | None,
|
|
1144
|
+
force: bool,
|
|
1145
|
+
) -> tuple[Path | None, str | None]:
|
|
1146
|
+
"""Decide where to write and which format to use.
|
|
1147
|
+
|
|
1148
|
+
Returns (output_path_or_None, format_name_or_None).
|
|
1149
|
+
output_path is None when no file output is configured.
|
|
1150
|
+
format_name is one of: csv, json, markdown - or None when output_path
|
|
1151
|
+
is None AND no --output-format was supplied (callers may treat
|
|
1152
|
+
that as "use the native display path").
|
|
1153
|
+
|
|
1154
|
+
Raises OutputFormatError for unknown extensions when --output-format
|
|
1155
|
+
isn't passed, and refuses to overwrite an existing file without --force.
|
|
1156
|
+
|
|
1157
|
+
Does NOT check input/output overlap - callers that have an input file
|
|
1158
|
+
must perform that check themselves.
|
|
1159
|
+
"""
|
|
1160
|
+
if output is None:
|
|
1161
|
+
if output_format:
|
|
1162
|
+
return None, output_format.lower()
|
|
1163
|
+
return None, None
|
|
1164
|
+
|
|
1165
|
+
output_path = Path(output).expanduser().resolve()
|
|
1166
|
+
|
|
1167
|
+
if output_path.exists() and not force:
|
|
1168
|
+
raise OutputFormatError(
|
|
1169
|
+
f"Output file already exists: {output_path}\n"
|
|
1170
|
+
f" Use --force to overwrite, or pick a different --output path."
|
|
1171
|
+
)
|
|
1172
|
+
|
|
1173
|
+
if output_format:
|
|
1174
|
+
fmt = output_format.lower()
|
|
1175
|
+
else:
|
|
1176
|
+
detected = detect_output_format(output_path)
|
|
1177
|
+
if detected is None:
|
|
1178
|
+
raise OutputFormatError(
|
|
1179
|
+
f"Cannot infer output format from {output_path.suffix!r}.\n"
|
|
1180
|
+
f" Pass --output-format csv|json|markdown explicitly, "
|
|
1181
|
+
f"or use a recognized extension (.csv .json .md)."
|
|
1182
|
+
)
|
|
1183
|
+
fmt = detected
|
|
1184
|
+
|
|
1185
|
+
# Ensure the parent directory exists - users sometimes pass
|
|
1186
|
+
# `./results/today.csv` without creating ./results/ first.
|
|
1187
|
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
1188
|
+
return output_path, fmt
|
|
1189
|
+
|
|
1190
|
+
|
|
1191
|
+
def _resolve_batch_output(
|
|
1192
|
+
*,
|
|
1193
|
+
input_path: str,
|
|
1194
|
+
output: str | None,
|
|
1195
|
+
output_format: str | None,
|
|
1196
|
+
force: bool,
|
|
1197
|
+
) -> tuple[Path | None, str]:
|
|
1198
|
+
"""Decide where to write and which format to use for the batch command.
|
|
1199
|
+
|
|
1200
|
+
Returns (output_path_or_None, format_name).
|
|
1201
|
+
output_path is None when writing to stdout.
|
|
1202
|
+
format_name is one of: csv, json, markdown (defaults to markdown
|
|
1203
|
+
for stdout when --output-format is not supplied).
|
|
1204
|
+
|
|
1205
|
+
Raises OutputFormatError for unknown extensions when --output-format
|
|
1206
|
+
isn't passed, refuses to overwrite an existing file without --force,
|
|
1207
|
+
and refuses to write output over the input file.
|
|
1208
|
+
"""
|
|
1209
|
+
# Overlap check runs BEFORE format detection so a user pointing --output
|
|
1210
|
+
# at an input file (no known output extension) sees the "input file"
|
|
1211
|
+
# error rather than the less-actionable "can't infer format" one.
|
|
1212
|
+
if output is not None:
|
|
1213
|
+
prospective = Path(output).expanduser().resolve()
|
|
1214
|
+
if output_overlaps_input(Path(input_path), prospective):
|
|
1215
|
+
raise OutputFormatError(
|
|
1216
|
+
f"Refusing to write output over the input file ({prospective}).\n"
|
|
1217
|
+
f" Choose a different --output path."
|
|
1218
|
+
)
|
|
1219
|
+
|
|
1220
|
+
output_path, fmt = _resolve_output_path(output, output_format, force)
|
|
1221
|
+
|
|
1222
|
+
if output_path is None:
|
|
1223
|
+
# Stdout default for batch: markdown.
|
|
1224
|
+
return None, (fmt or "markdown")
|
|
1225
|
+
|
|
1226
|
+
assert fmt is not None # _resolve_output_path guarantees this when path is non-None
|
|
1227
|
+
return output_path, fmt
|
|
1228
|
+
|
|
1229
|
+
|
|
1230
|
+
async def _run_mode_only_judging(
|
|
1231
|
+
*,
|
|
1232
|
+
states: list[StreamState],
|
|
1233
|
+
prompt: str,
|
|
1234
|
+
judge_models: list[str],
|
|
1235
|
+
criteria: list[str] | None,
|
|
1236
|
+
template: str,
|
|
1237
|
+
provider_factory: Callable[[str], BaseProvider],
|
|
1238
|
+
concurrency: int,
|
|
1239
|
+
) -> list[JudgeResult]:
|
|
1240
|
+
"""Judge only the mode output per cell, then expand the verdict to every run.
|
|
1241
|
+
|
|
1242
|
+
With --runs N, judging every run is expensive. We pick one canonical
|
|
1243
|
+
representative per (model, temperature, system_prompt) cell - the
|
|
1244
|
+
mode output when a clear winner exists, otherwise the first
|
|
1245
|
+
successful run - and assign that single verdict to every state in
|
|
1246
|
+
the cell. Returns a list of JudgeResult parallel to `states`.
|
|
1247
|
+
"""
|
|
1248
|
+
from cli_modelarium.run_statistics import compute_run_stats, group_states_by_cell
|
|
1249
|
+
|
|
1250
|
+
groups = group_states_by_cell(states)
|
|
1251
|
+
|
|
1252
|
+
# Pick one representative state per cell.
|
|
1253
|
+
representatives: list[tuple[tuple[str, float, str | None], StreamState]] = []
|
|
1254
|
+
for key, cell_states in groups.items():
|
|
1255
|
+
stats = compute_run_stats(cell_states)
|
|
1256
|
+
chosen: StreamState | None = None
|
|
1257
|
+
if stats.mode_output is not None:
|
|
1258
|
+
for s in cell_states:
|
|
1259
|
+
if s.error is None and s.text == stats.mode_output:
|
|
1260
|
+
chosen = s
|
|
1261
|
+
break
|
|
1262
|
+
if chosen is None:
|
|
1263
|
+
# No mode (all unique) or all failed: fall back to first
|
|
1264
|
+
# successful state; if none, the cell stays unjudged.
|
|
1265
|
+
for s in cell_states:
|
|
1266
|
+
if s.error is None:
|
|
1267
|
+
chosen = s
|
|
1268
|
+
break
|
|
1269
|
+
if chosen is not None:
|
|
1270
|
+
representatives.append((key, chosen))
|
|
1271
|
+
|
|
1272
|
+
if not representatives:
|
|
1273
|
+
return [JudgeResult() for _ in states]
|
|
1274
|
+
|
|
1275
|
+
cell_verdicts = await run_judging(
|
|
1276
|
+
items=[(s, prompt) for _, s in representatives],
|
|
1277
|
+
judge_models=judge_models,
|
|
1278
|
+
criteria=criteria,
|
|
1279
|
+
provider_factory=provider_factory,
|
|
1280
|
+
template=template,
|
|
1281
|
+
response_parser=None,
|
|
1282
|
+
skip_self_eval=True,
|
|
1283
|
+
concurrency=concurrency,
|
|
1284
|
+
)
|
|
1285
|
+
|
|
1286
|
+
verdict_by_cell: dict[tuple[str, float, str | None], JudgeResult] = {
|
|
1287
|
+
key: verdict for (key, _), verdict in zip(representatives, cell_verdicts, strict=True)
|
|
1288
|
+
}
|
|
1289
|
+
|
|
1290
|
+
# Expand: every state in a cell gets that cell's single verdict.
|
|
1291
|
+
return [
|
|
1292
|
+
verdict_by_cell.get(
|
|
1293
|
+
(s.model, s.temperature, s.system_prompt),
|
|
1294
|
+
JudgeResult(),
|
|
1295
|
+
)
|
|
1296
|
+
for s in states
|
|
1297
|
+
]
|
|
1298
|
+
|
|
1299
|
+
|
|
1300
|
+
def _states_to_compare_results(
|
|
1301
|
+
states: list[StreamState],
|
|
1302
|
+
prompt: str,
|
|
1303
|
+
judge_results: list[JudgeResult] | None = None,
|
|
1304
|
+
) -> list[BatchResult]:
|
|
1305
|
+
"""Convert compare's flat StreamState list to BatchResult shape.
|
|
1306
|
+
|
|
1307
|
+
Each state becomes a BatchResult with a synthetic BatchPrompt
|
|
1308
|
+
(id=p1, p2, ... matching batch's auto-id convention from
|
|
1309
|
+
`batch._parse_txt`) so that the existing batch formatters can
|
|
1310
|
+
serialize compare results without a parallel codepath.
|
|
1311
|
+
"""
|
|
1312
|
+
results: list[BatchResult] = []
|
|
1313
|
+
for i, state in enumerate(states):
|
|
1314
|
+
bp = BatchPrompt(
|
|
1315
|
+
id=f"p{i + 1}",
|
|
1316
|
+
prompt=prompt,
|
|
1317
|
+
system=state.system_prompt,
|
|
1318
|
+
)
|
|
1319
|
+
jr = judge_results[i] if judge_results is not None else None
|
|
1320
|
+
results.append(state_to_result(state, bp, judge_result=jr, assertion_results=None))
|
|
1321
|
+
return results
|
|
1322
|
+
|
|
1323
|
+
|
|
1324
|
+
def _flatten_cell_cis(
|
|
1325
|
+
cis: dict,
|
|
1326
|
+
) -> dict:
|
|
1327
|
+
"""Convert {model: {metric: ConfidenceInterval}} to a flat dict keyed by
|
|
1328
|
+
model for formatter consumption.
|
|
1329
|
+
|
|
1330
|
+
Output shape (per model):
|
|
1331
|
+
{model_name: {
|
|
1332
|
+
"latency_ms": {"ci_low": .., "ci_high": .., "ci_level": ..,
|
|
1333
|
+
"method": .., "n_resamples": .., "seed": ..},
|
|
1334
|
+
...
|
|
1335
|
+
}}
|
|
1336
|
+
"""
|
|
1337
|
+
out: dict = {}
|
|
1338
|
+
for model, metrics in cis.items():
|
|
1339
|
+
out[model] = {}
|
|
1340
|
+
for metric, ci in metrics.items():
|
|
1341
|
+
if ci is None:
|
|
1342
|
+
continue
|
|
1343
|
+
out[model][metric] = {
|
|
1344
|
+
"ci_low": ci.ci_low,
|
|
1345
|
+
"ci_high": ci.ci_high,
|
|
1346
|
+
"ci_level": ci.ci_level,
|
|
1347
|
+
"method": ci.method,
|
|
1348
|
+
"n_resamples": ci.n_resamples,
|
|
1349
|
+
"seed": ci.seed,
|
|
1350
|
+
}
|
|
1351
|
+
return out
|
|
1352
|
+
|
|
1353
|
+
|
|
1354
|
+
def _emit_batch_results(
|
|
1355
|
+
results: list,
|
|
1356
|
+
*,
|
|
1357
|
+
output_path: Path | None,
|
|
1358
|
+
output_fmt: str,
|
|
1359
|
+
runs: int = 1,
|
|
1360
|
+
significance_results: list | None = None,
|
|
1361
|
+
stats_by_cell_cis: dict | None = None,
|
|
1362
|
+
mcnemar_results: list | None = None,
|
|
1363
|
+
methodology: dict | None = None,
|
|
1364
|
+
) -> None:
|
|
1365
|
+
"""Dispatch to the right writer/renderer based on resolved format.
|
|
1366
|
+
|
|
1367
|
+
v0.1.3: significance_results, stats_by_cell_cis, mcnemar_results, and
|
|
1368
|
+
methodology flow to ALL formatters (not just JSON) so CSV/Markdown
|
|
1369
|
+
also render the new fields.
|
|
1370
|
+
"""
|
|
1371
|
+
if output_path is None:
|
|
1372
|
+
# Stdout: only markdown is rendered natively; csv/json get printed raw.
|
|
1373
|
+
if output_fmt == "markdown":
|
|
1374
|
+
render_markdown_to_console(
|
|
1375
|
+
results,
|
|
1376
|
+
console,
|
|
1377
|
+
runs=runs,
|
|
1378
|
+
significance_results=significance_results,
|
|
1379
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1380
|
+
mcnemar_results=mcnemar_results,
|
|
1381
|
+
methodology=methodology,
|
|
1382
|
+
)
|
|
1383
|
+
elif output_fmt == "csv":
|
|
1384
|
+
from cli_modelarium.output_formatters import _format_csv
|
|
1385
|
+
|
|
1386
|
+
console.print(
|
|
1387
|
+
_format_csv(
|
|
1388
|
+
results,
|
|
1389
|
+
runs=runs,
|
|
1390
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1391
|
+
),
|
|
1392
|
+
end="",
|
|
1393
|
+
)
|
|
1394
|
+
elif output_fmt == "json":
|
|
1395
|
+
from cli_modelarium.output_formatters import _format_json
|
|
1396
|
+
|
|
1397
|
+
console.print(
|
|
1398
|
+
_format_json(
|
|
1399
|
+
results,
|
|
1400
|
+
runs=runs,
|
|
1401
|
+
significance_results=significance_results,
|
|
1402
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1403
|
+
mcnemar_results=mcnemar_results,
|
|
1404
|
+
methodology=methodology,
|
|
1405
|
+
),
|
|
1406
|
+
end="",
|
|
1407
|
+
)
|
|
1408
|
+
else:
|
|
1409
|
+
_print_error(f"Unsupported output format: {output_fmt!r}")
|
|
1410
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1411
|
+
return
|
|
1412
|
+
|
|
1413
|
+
if output_fmt == "csv":
|
|
1414
|
+
write_csv(
|
|
1415
|
+
results,
|
|
1416
|
+
output_path,
|
|
1417
|
+
runs=runs,
|
|
1418
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1419
|
+
)
|
|
1420
|
+
elif output_fmt == "json":
|
|
1421
|
+
write_json(
|
|
1422
|
+
results,
|
|
1423
|
+
output_path,
|
|
1424
|
+
runs=runs,
|
|
1425
|
+
significance_results=significance_results,
|
|
1426
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1427
|
+
mcnemar_results=mcnemar_results,
|
|
1428
|
+
methodology=methodology,
|
|
1429
|
+
)
|
|
1430
|
+
elif output_fmt == "markdown":
|
|
1431
|
+
write_markdown(
|
|
1432
|
+
results,
|
|
1433
|
+
output_path,
|
|
1434
|
+
runs=runs,
|
|
1435
|
+
significance_results=significance_results,
|
|
1436
|
+
stats_by_cell_cis=stats_by_cell_cis,
|
|
1437
|
+
mcnemar_results=mcnemar_results,
|
|
1438
|
+
methodology=methodology,
|
|
1439
|
+
)
|
|
1440
|
+
else:
|
|
1441
|
+
_print_error(f"Unsupported output format: {output_fmt!r}")
|
|
1442
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1443
|
+
console.print(f"[dim]Wrote {output_path}[/dim]")
|
|
1444
|
+
|
|
1445
|
+
|
|
1446
|
+
# ===== configure =====
|
|
1447
|
+
|
|
1448
|
+
|
|
1449
|
+
@main.command()
|
|
1450
|
+
def configure() -> None:
|
|
1451
|
+
"""Interactively set API keys for each provider."""
|
|
1452
|
+
providers = [p for p in all_known_providers() if p != "local"]
|
|
1453
|
+
|
|
1454
|
+
console.print(
|
|
1455
|
+
Panel(
|
|
1456
|
+
"Configure API keys. Keys are stored in your OS-native keychain.\n"
|
|
1457
|
+
"Press Enter to skip any provider.",
|
|
1458
|
+
title="cli-modelarium setup",
|
|
1459
|
+
border_style="cyan",
|
|
1460
|
+
)
|
|
1461
|
+
)
|
|
1462
|
+
|
|
1463
|
+
saved = 0
|
|
1464
|
+
for provider in providers:
|
|
1465
|
+
try:
|
|
1466
|
+
key = Prompt.ask(
|
|
1467
|
+
f"{provider.capitalize()} API key",
|
|
1468
|
+
password=True,
|
|
1469
|
+
default="",
|
|
1470
|
+
show_default=False,
|
|
1471
|
+
)
|
|
1472
|
+
except (EOFError, KeyboardInterrupt):
|
|
1473
|
+
console.print("\n[yellow]Setup cancelled.[/yellow]")
|
|
1474
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1475
|
+
|
|
1476
|
+
if not key.strip():
|
|
1477
|
+
console.print(f" [dim]Skipped {provider}[/dim]")
|
|
1478
|
+
continue
|
|
1479
|
+
|
|
1480
|
+
try:
|
|
1481
|
+
save_key(provider, key)
|
|
1482
|
+
except ValueError as e:
|
|
1483
|
+
console.print(f" [red]Invalid format - {e}[/red]")
|
|
1484
|
+
continue
|
|
1485
|
+
except Exception as e:
|
|
1486
|
+
console.print(f" [red]Could not save: {redact_secrets(str(e))}[/red]")
|
|
1487
|
+
continue
|
|
1488
|
+
|
|
1489
|
+
saved += 1
|
|
1490
|
+
console.print(f" [green]Saved {provider} to keychain[/green]")
|
|
1491
|
+
|
|
1492
|
+
console.print()
|
|
1493
|
+
console.print(
|
|
1494
|
+
Panel(
|
|
1495
|
+
f"{saved} of {len(providers)} providers configured.\nRun: cli-modelarium list-models",
|
|
1496
|
+
title="Configuration complete",
|
|
1497
|
+
border_style="green",
|
|
1498
|
+
)
|
|
1499
|
+
)
|
|
1500
|
+
|
|
1501
|
+
|
|
1502
|
+
# ===== keys =====
|
|
1503
|
+
|
|
1504
|
+
|
|
1505
|
+
@main.group()
|
|
1506
|
+
def keys() -> None:
|
|
1507
|
+
"""Manage API keys (stored in OS-native keychain)."""
|
|
1508
|
+
|
|
1509
|
+
|
|
1510
|
+
@keys.command("list")
|
|
1511
|
+
def keys_list() -> None:
|
|
1512
|
+
"""Show which providers have keys configured."""
|
|
1513
|
+
providers = [p for p in all_known_providers() if p != "local"]
|
|
1514
|
+
|
|
1515
|
+
table = Table(title="API key status", border_style="dim")
|
|
1516
|
+
table.add_column("Provider", style="bold")
|
|
1517
|
+
table.add_column("Status")
|
|
1518
|
+
|
|
1519
|
+
for provider in providers:
|
|
1520
|
+
if is_key_configured(provider):
|
|
1521
|
+
table.add_row(provider, "[green]configured[/green]")
|
|
1522
|
+
else:
|
|
1523
|
+
table.add_row(provider, "[dim]not configured[/dim]")
|
|
1524
|
+
|
|
1525
|
+
saved_local = load_local_url()
|
|
1526
|
+
if saved_local:
|
|
1527
|
+
table.add_row("local", f"[green]{saved_local}[/green]")
|
|
1528
|
+
else:
|
|
1529
|
+
table.add_row("local", f"[dim]default ({LocalProvider.DEFAULT_URL})[/dim]")
|
|
1530
|
+
|
|
1531
|
+
console.print(table)
|
|
1532
|
+
|
|
1533
|
+
|
|
1534
|
+
@keys.command("set")
|
|
1535
|
+
@click.argument("provider")
|
|
1536
|
+
@click.option("--base-url", help="(local provider only) Override default base URL.")
|
|
1537
|
+
def keys_set(provider: str, base_url: str | None) -> None:
|
|
1538
|
+
"""Set or update the API key for a provider (prompts securely).
|
|
1539
|
+
|
|
1540
|
+
For the local provider, pass --base-url to persist a default URL
|
|
1541
|
+
instead of prompting for an API key.
|
|
1542
|
+
"""
|
|
1543
|
+
if provider == "local":
|
|
1544
|
+
if not base_url:
|
|
1545
|
+
_print_error(
|
|
1546
|
+
"Local provider takes --base-url, not an API key.\n"
|
|
1547
|
+
" Example: cli-modelarium keys set local --base-url http://localhost:1234/v1"
|
|
1548
|
+
)
|
|
1549
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1550
|
+
try:
|
|
1551
|
+
LocalProvider._validate_local_url(base_url)
|
|
1552
|
+
except ModelariumError as e:
|
|
1553
|
+
_print_error(str(e))
|
|
1554
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1555
|
+
save_local_url(base_url)
|
|
1556
|
+
console.print(f"[green]Saved local provider URL: {base_url}[/green]")
|
|
1557
|
+
return
|
|
1558
|
+
|
|
1559
|
+
if provider not in KEY_PATTERNS:
|
|
1560
|
+
_print_error(
|
|
1561
|
+
f"Unknown provider: {provider}.\n"
|
|
1562
|
+
f"Supported providers: {', '.join(sorted(KEY_PATTERNS))}, local"
|
|
1563
|
+
)
|
|
1564
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1565
|
+
|
|
1566
|
+
try:
|
|
1567
|
+
key = Prompt.ask(f"{provider.capitalize()} API key", password=True)
|
|
1568
|
+
except (EOFError, KeyboardInterrupt):
|
|
1569
|
+
console.print("\n[yellow]Cancelled.[/yellow]")
|
|
1570
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1571
|
+
|
|
1572
|
+
try:
|
|
1573
|
+
save_key(provider, key)
|
|
1574
|
+
except ValueError as e:
|
|
1575
|
+
_print_error(str(e))
|
|
1576
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1577
|
+
except Exception as e:
|
|
1578
|
+
_print_error(redact_secrets(str(e)))
|
|
1579
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1580
|
+
|
|
1581
|
+
console.print(f"[green]Saved {provider} key to keychain.[/green]")
|
|
1582
|
+
|
|
1583
|
+
|
|
1584
|
+
@keys.command("delete")
|
|
1585
|
+
@click.argument("provider")
|
|
1586
|
+
def keys_delete(provider: str) -> None:
|
|
1587
|
+
"""Remove the API key for a provider from the keychain."""
|
|
1588
|
+
if provider != "local" and provider not in KEY_PATTERNS:
|
|
1589
|
+
_print_error(
|
|
1590
|
+
f"Unknown provider: {provider}.\n"
|
|
1591
|
+
f"Supported providers: {', '.join(sorted(KEY_PATTERNS))}, local"
|
|
1592
|
+
)
|
|
1593
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1594
|
+
|
|
1595
|
+
if provider == "local":
|
|
1596
|
+
if delete_local_url():
|
|
1597
|
+
console.print("[green]Removed saved local provider URL.[/green]")
|
|
1598
|
+
else:
|
|
1599
|
+
console.print("[dim]No saved local provider URL.[/dim]")
|
|
1600
|
+
return
|
|
1601
|
+
|
|
1602
|
+
if delete_key(provider):
|
|
1603
|
+
console.print(f"[green]Removed {provider} key from keychain.[/green]")
|
|
1604
|
+
else:
|
|
1605
|
+
console.print(f"[dim]No {provider} key was stored.[/dim]")
|
|
1606
|
+
|
|
1607
|
+
|
|
1608
|
+
# ===== list-models =====
|
|
1609
|
+
|
|
1610
|
+
|
|
1611
|
+
@main.command("list-models")
|
|
1612
|
+
@click.option(
|
|
1613
|
+
"--local", "local_only", is_flag=True, help="Show only local models (queries the local server)."
|
|
1614
|
+
)
|
|
1615
|
+
@click.option("--local-url", help="Override default URL for local-model discovery.")
|
|
1616
|
+
def list_models(local_only: bool, local_url: str | None) -> None:
|
|
1617
|
+
"""List supported models, grouped by provider."""
|
|
1618
|
+
if local_only:
|
|
1619
|
+
_list_local_models(local_url)
|
|
1620
|
+
return
|
|
1621
|
+
|
|
1622
|
+
providers = all_known_providers()
|
|
1623
|
+
|
|
1624
|
+
any_shown = False
|
|
1625
|
+
for provider in providers:
|
|
1626
|
+
models = list_models_for_provider(provider)
|
|
1627
|
+
if provider == "local":
|
|
1628
|
+
# Local models are dynamic - skip the static section; we show
|
|
1629
|
+
# discovered models when --local is passed.
|
|
1630
|
+
continue
|
|
1631
|
+
if not models:
|
|
1632
|
+
continue
|
|
1633
|
+
|
|
1634
|
+
any_shown = True
|
|
1635
|
+
configured = "configured" if is_key_configured(provider) else "not configured"
|
|
1636
|
+
title = f"{provider} [dim]({configured})[/dim]"
|
|
1637
|
+
|
|
1638
|
+
table = Table(title=title, border_style="dim", title_justify="left")
|
|
1639
|
+
table.add_column("Model", style="bold")
|
|
1640
|
+
table.add_column("Input $/MTok", justify="right")
|
|
1641
|
+
table.add_column("Output $/MTok", justify="right")
|
|
1642
|
+
table.add_column("Cached $/MTok", justify="right", style="dim")
|
|
1643
|
+
|
|
1644
|
+
for model in models:
|
|
1645
|
+
entry = PRICING[model]
|
|
1646
|
+
cached = entry.get("cached_input")
|
|
1647
|
+
cached_text = f"${float(cached):.4f}" if cached is not None else "-"
|
|
1648
|
+
table.add_row(
|
|
1649
|
+
model,
|
|
1650
|
+
f"${float(entry['input']):.4f}",
|
|
1651
|
+
f"${float(entry['output']):.4f}",
|
|
1652
|
+
cached_text,
|
|
1653
|
+
)
|
|
1654
|
+
|
|
1655
|
+
console.print(table)
|
|
1656
|
+
console.print()
|
|
1657
|
+
|
|
1658
|
+
if not any_shown:
|
|
1659
|
+
_print_error("No models registered.")
|
|
1660
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1661
|
+
|
|
1662
|
+
console.print(
|
|
1663
|
+
"[dim]Local models are routed by the `local/` prefix.[/dim] "
|
|
1664
|
+
"[dim]Run `cli-modelarium list-models --local` to discover what's running locally.[/dim]"
|
|
1665
|
+
)
|
|
1666
|
+
console.print(f"[dim]{pricing_freshness_note()}[/dim]")
|
|
1667
|
+
|
|
1668
|
+
|
|
1669
|
+
def _list_local_models(local_url: str | None) -> None:
|
|
1670
|
+
"""Query the local server's /models endpoint and render the result."""
|
|
1671
|
+
url = local_url or load_local_url() or LocalProvider.DEFAULT_URL
|
|
1672
|
+
|
|
1673
|
+
try:
|
|
1674
|
+
# Validate the URL up front - we want LocalURLError BEFORE any I/O.
|
|
1675
|
+
LocalProvider._validate_local_url(url)
|
|
1676
|
+
except ModelariumError as e:
|
|
1677
|
+
_print_error(str(e))
|
|
1678
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1679
|
+
|
|
1680
|
+
try:
|
|
1681
|
+
models = asyncio.run(LocalProvider.discover_models(url))
|
|
1682
|
+
except httpx.ConnectError:
|
|
1683
|
+
console.print(
|
|
1684
|
+
Panel(
|
|
1685
|
+
f"Could not reach local server at {url}.\n\n"
|
|
1686
|
+
f"Possible causes:\n"
|
|
1687
|
+
f" - Server not running (try: ollama serve)\n"
|
|
1688
|
+
f" - Wrong URL (use --local-url to override)\n"
|
|
1689
|
+
f" - Firewall blocking the connection",
|
|
1690
|
+
title="Local models",
|
|
1691
|
+
border_style="yellow",
|
|
1692
|
+
)
|
|
1693
|
+
)
|
|
1694
|
+
return
|
|
1695
|
+
except httpx.TimeoutException:
|
|
1696
|
+
console.print(
|
|
1697
|
+
Panel(
|
|
1698
|
+
f"Timed out connecting to {url} "
|
|
1699
|
+
f"(waited {LocalProvider.DISCOVERY_TIMEOUT_SECONDS:.0f}s).\n"
|
|
1700
|
+
f"The server may be starting up or under heavy load.",
|
|
1701
|
+
title="Local models",
|
|
1702
|
+
border_style="yellow",
|
|
1703
|
+
)
|
|
1704
|
+
)
|
|
1705
|
+
return
|
|
1706
|
+
except (httpx.HTTPStatusError, httpx.RequestError, ValueError) as e:
|
|
1707
|
+
console.print(
|
|
1708
|
+
Panel(
|
|
1709
|
+
f"Local server at {url} returned an unexpected response:\n"
|
|
1710
|
+
f" {redact_secrets(str(e))}",
|
|
1711
|
+
title="Local models",
|
|
1712
|
+
border_style="yellow",
|
|
1713
|
+
)
|
|
1714
|
+
)
|
|
1715
|
+
return
|
|
1716
|
+
|
|
1717
|
+
if not models:
|
|
1718
|
+
console.print(
|
|
1719
|
+
Panel(
|
|
1720
|
+
f"Local server at {url} responded but has no models installed.\n"
|
|
1721
|
+
f"For Ollama: ollama pull llama3.3",
|
|
1722
|
+
title="Local models",
|
|
1723
|
+
border_style="cyan",
|
|
1724
|
+
)
|
|
1725
|
+
)
|
|
1726
|
+
return
|
|
1727
|
+
|
|
1728
|
+
table = Table(
|
|
1729
|
+
title=f"local [dim]({url})[/dim]",
|
|
1730
|
+
border_style="dim",
|
|
1731
|
+
title_justify="left",
|
|
1732
|
+
)
|
|
1733
|
+
table.add_column("Model ID for cli-modelarium", style="bold")
|
|
1734
|
+
table.add_column("Created", style="dim")
|
|
1735
|
+
table.add_column("Owned by", style="dim")
|
|
1736
|
+
table.add_column("Cost", justify="right")
|
|
1737
|
+
|
|
1738
|
+
for entry in models:
|
|
1739
|
+
model_id = entry.get("id", "(unnamed)")
|
|
1740
|
+
created = entry.get("created")
|
|
1741
|
+
created_text = _format_unix_timestamp(created) if created else "-"
|
|
1742
|
+
owned_by = str(entry.get("owned_by", "-"))
|
|
1743
|
+
table.add_row(f"local/{model_id}", created_text, owned_by, "[dim]Free[/dim]")
|
|
1744
|
+
|
|
1745
|
+
console.print(table)
|
|
1746
|
+
first_id = models[0].get("id", "<name>")
|
|
1747
|
+
console.print(f"\n[dim]Use these via: cli-modelarium 'prompt' --models local/{first_id}[/dim]")
|
|
1748
|
+
|
|
1749
|
+
|
|
1750
|
+
def _format_unix_timestamp(ts: object) -> str:
|
|
1751
|
+
"""Format a unix timestamp (int or float) as YYYY-MM-DD, or '-' if unparseable."""
|
|
1752
|
+
try:
|
|
1753
|
+
from datetime import UTC, datetime
|
|
1754
|
+
|
|
1755
|
+
return datetime.fromtimestamp(float(ts), tz=UTC).strftime("%Y-%m-%d")
|
|
1756
|
+
except (TypeError, ValueError, OSError):
|
|
1757
|
+
return "-"
|
|
1758
|
+
|
|
1759
|
+
|
|
1760
|
+
# ===== pricing =====
|
|
1761
|
+
|
|
1762
|
+
|
|
1763
|
+
@main.command("pricing")
|
|
1764
|
+
@click.argument("model", required=False)
|
|
1765
|
+
@click.option("--all", "show_all", is_flag=True, help="Show pricing for every model.")
|
|
1766
|
+
def pricing_cmd(model: str | None, show_all: bool) -> None:
|
|
1767
|
+
"""Show pricing for a model or all models."""
|
|
1768
|
+
if show_all or model is None:
|
|
1769
|
+
table = Table(title="Pricing (per 1M tokens, USD)", border_style="dim")
|
|
1770
|
+
table.add_column("Model", style="bold")
|
|
1771
|
+
table.add_column("Provider")
|
|
1772
|
+
table.add_column("Input", justify="right")
|
|
1773
|
+
table.add_column("Output", justify="right")
|
|
1774
|
+
table.add_column("Cached", justify="right", style="dim")
|
|
1775
|
+
|
|
1776
|
+
for name in sorted(PRICING):
|
|
1777
|
+
if name.endswith("/*"):
|
|
1778
|
+
continue
|
|
1779
|
+
entry = PRICING[name]
|
|
1780
|
+
if entry.get("is_local"):
|
|
1781
|
+
table.add_row(
|
|
1782
|
+
name,
|
|
1783
|
+
str(entry["provider"]),
|
|
1784
|
+
"[dim]Free[/dim]",
|
|
1785
|
+
"[dim]Free[/dim]",
|
|
1786
|
+
"[dim]-[/dim]",
|
|
1787
|
+
)
|
|
1788
|
+
continue
|
|
1789
|
+
cached = entry.get("cached_input")
|
|
1790
|
+
cached_text = f"${float(cached):.4f}" if cached is not None else "-"
|
|
1791
|
+
table.add_row(
|
|
1792
|
+
name,
|
|
1793
|
+
str(entry["provider"]),
|
|
1794
|
+
f"${float(entry['input']):.4f}",
|
|
1795
|
+
f"${float(entry['output']):.4f}",
|
|
1796
|
+
cached_text,
|
|
1797
|
+
)
|
|
1798
|
+
|
|
1799
|
+
console.print(table)
|
|
1800
|
+
console.print(f"[dim]{pricing_freshness_note()}[/dim]")
|
|
1801
|
+
return
|
|
1802
|
+
|
|
1803
|
+
if is_local_model(model):
|
|
1804
|
+
console.print(f"[bold]{model}[/bold]: [dim]Free (local model)[/dim]")
|
|
1805
|
+
return
|
|
1806
|
+
|
|
1807
|
+
entry = PRICING.get(model)
|
|
1808
|
+
if entry is None:
|
|
1809
|
+
_print_error(f"Unknown model: {model}. Run `cli-modelarium list-models` to see options.")
|
|
1810
|
+
sys.exit(EXIT_CALL_FAILED)
|
|
1811
|
+
|
|
1812
|
+
cached = entry.get("cached_input")
|
|
1813
|
+
console.print(f"[bold]{model}[/bold] ([dim]{entry['provider']}[/dim])")
|
|
1814
|
+
console.print(f" Input: ${float(entry['input']):.4f} / 1M tokens")
|
|
1815
|
+
console.print(f" Output: ${float(entry['output']):.4f} / 1M tokens")
|
|
1816
|
+
if cached is not None:
|
|
1817
|
+
console.print(f" Cached: ${float(cached):.4f} / 1M tokens")
|
|
1818
|
+
console.print(f"\n[dim]{pricing_freshness_note()}[/dim]")
|
|
1819
|
+
|
|
1820
|
+
|
|
1821
|
+
# ===== helpers =====
|
|
1822
|
+
|
|
1823
|
+
|
|
1824
|
+
def _parse_temperatures(raw: str) -> list[float]:
|
|
1825
|
+
"""Parse a comma-separated temperatures string into floats. Defaults to [0.0]."""
|
|
1826
|
+
if not raw.strip():
|
|
1827
|
+
return [0.0]
|
|
1828
|
+
out: list[float] = []
|
|
1829
|
+
for token in raw.split(","):
|
|
1830
|
+
token = token.strip()
|
|
1831
|
+
if not token:
|
|
1832
|
+
continue
|
|
1833
|
+
try:
|
|
1834
|
+
out.append(float(token))
|
|
1835
|
+
except ValueError:
|
|
1836
|
+
raise ValueError(f"Invalid temperature value: {token!r}") from None
|
|
1837
|
+
return out or [0.0]
|
|
1838
|
+
|
|
1839
|
+
|
|
1840
|
+
def _resolve_system_prompts(
|
|
1841
|
+
*,
|
|
1842
|
+
system_prompt: str | None,
|
|
1843
|
+
system_prompts: str | None,
|
|
1844
|
+
system_prompt_file: str | None,
|
|
1845
|
+
) -> list[str | None]:
|
|
1846
|
+
"""Resolve the three mutually-exclusive system-prompt flags into a list.
|
|
1847
|
+
|
|
1848
|
+
Returns `[None]` when no system prompt is configured (the orchestrator
|
|
1849
|
+
treats this as "one task with no system prompt"). Returns a list of
|
|
1850
|
+
strings otherwise - never an empty list.
|
|
1851
|
+
|
|
1852
|
+
Raises:
|
|
1853
|
+
click.UsageError: if more than one of the three flags is set.
|
|
1854
|
+
FileNotFoundError / ValueError: from `load_system_prompt`.
|
|
1855
|
+
"""
|
|
1856
|
+
used = [
|
|
1857
|
+
name
|
|
1858
|
+
for name, val in (
|
|
1859
|
+
("--system-prompt", system_prompt),
|
|
1860
|
+
("--system-prompts", system_prompts),
|
|
1861
|
+
("--system-prompt-file", system_prompt_file),
|
|
1862
|
+
)
|
|
1863
|
+
if val
|
|
1864
|
+
]
|
|
1865
|
+
if len(used) > 1:
|
|
1866
|
+
raise click.UsageError(f"{', '.join(used)} are mutually exclusive - pick one.")
|
|
1867
|
+
|
|
1868
|
+
if system_prompt_file:
|
|
1869
|
+
return [load_system_prompt(system_prompt_file)]
|
|
1870
|
+
if system_prompts:
|
|
1871
|
+
parsed = _split_system_prompts(system_prompts)
|
|
1872
|
+
return parsed or [None]
|
|
1873
|
+
if system_prompt:
|
|
1874
|
+
# Empty-string case is also caught here; we'd have failed the `if`
|
|
1875
|
+
# above. But guard anyway: a stripped-empty value means no prompt.
|
|
1876
|
+
stripped = system_prompt.strip()
|
|
1877
|
+
return [stripped] if stripped else [None]
|
|
1878
|
+
return [None]
|
|
1879
|
+
|
|
1880
|
+
|
|
1881
|
+
# Re-export under the cli.py namespace for backward compatibility with
|
|
1882
|
+
# Phase 6 and Phase 8 tests that import these private aliases directly.
|
|
1883
|
+
_split_escaped_csv = split_escaped_csv
|
|
1884
|
+
_split_system_prompts = split_escaped_csv
|
|
1885
|
+
|
|
1886
|
+
|
|
1887
|
+
def _resolve_judge_models(*, judge: str | None, judges: str | None) -> list[str]:
|
|
1888
|
+
"""Resolve --judge / --judges into a list of judge model IDs.
|
|
1889
|
+
|
|
1890
|
+
Returns [] when neither flag is set. Raises click.UsageError if both
|
|
1891
|
+
flags are set simultaneously (they're mutually exclusive).
|
|
1892
|
+
"""
|
|
1893
|
+
if judge and judges:
|
|
1894
|
+
raise click.UsageError("--judge and --judges are mutually exclusive - pick one.")
|
|
1895
|
+
if judges:
|
|
1896
|
+
return _split_escaped_csv(judges)
|
|
1897
|
+
if judge and judge.strip():
|
|
1898
|
+
return [judge.strip()]
|
|
1899
|
+
return []
|
|
1900
|
+
|
|
1901
|
+
|
|
1902
|
+
def _resolve_judge_criteria_and_template(
|
|
1903
|
+
*,
|
|
1904
|
+
judge_criteria: str | None,
|
|
1905
|
+
judge_template: str | None,
|
|
1906
|
+
) -> tuple[list[str], str]:
|
|
1907
|
+
"""Resolve --judge-criteria and --judge-template into (criteria, template).
|
|
1908
|
+
|
|
1909
|
+
Returns (DEFAULT_CRITERIA, JUDGE_PROMPT_TEMPLATE) when neither is set.
|
|
1910
|
+
Raises click.UsageError if both are set simultaneously.
|
|
1911
|
+
|
|
1912
|
+
--judge-template loads a custom prompt template from disk (UTF-8, max 1 MB).
|
|
1913
|
+
--judge-criteria splits on commas with the same `\\,` escape as system prompts.
|
|
1914
|
+
"""
|
|
1915
|
+
if judge_criteria and judge_template:
|
|
1916
|
+
raise click.UsageError(
|
|
1917
|
+
"--judge-criteria and --judge-template are mutually exclusive - pick one."
|
|
1918
|
+
)
|
|
1919
|
+
criteria = list(DEFAULT_CRITERIA)
|
|
1920
|
+
template = JUDGE_PROMPT_TEMPLATE
|
|
1921
|
+
if judge_criteria:
|
|
1922
|
+
criteria = _split_escaped_csv(judge_criteria) or list(DEFAULT_CRITERIA)
|
|
1923
|
+
if judge_template:
|
|
1924
|
+
# Reuse the system-prompt-file loader: same size + encoding contract.
|
|
1925
|
+
template = load_system_prompt(judge_template)
|
|
1926
|
+
return criteria, template
|
|
1927
|
+
|
|
1928
|
+
|
|
1929
|
+
def _validate_judge_models(judge_models: list[str], *, local_url: str | None) -> None:
|
|
1930
|
+
"""Ensure every judge model is in the registry AND has a configured key.
|
|
1931
|
+
|
|
1932
|
+
This runs BEFORE any main API calls - the build prompt's contract is
|
|
1933
|
+
that a misconfigured judge fails the run immediately, not after burning
|
|
1934
|
+
money on the comparison.
|
|
1935
|
+
"""
|
|
1936
|
+
from cli_modelarium.models_registry import get_provider_for_model
|
|
1937
|
+
from cli_modelarium.security import is_key_configured
|
|
1938
|
+
|
|
1939
|
+
seen_providers: set[str] = set()
|
|
1940
|
+
for model in judge_models:
|
|
1941
|
+
provider_name = get_provider_for_model(model) # raises UnknownModelError
|
|
1942
|
+
if provider_name in seen_providers:
|
|
1943
|
+
continue
|
|
1944
|
+
seen_providers.add(provider_name)
|
|
1945
|
+
if provider_name == "local":
|
|
1946
|
+
# Local needs no key; the URL check happens at construction time.
|
|
1947
|
+
continue
|
|
1948
|
+
if not is_key_configured(provider_name):
|
|
1949
|
+
raise KeyNotConfiguredError(provider_name)
|
|
1950
|
+
|
|
1951
|
+
|
|
1952
|
+
def _get_provider_instance(provider_name: str, *, local_url: str | None = None) -> BaseProvider:
|
|
1953
|
+
"""Instantiate the provider for `provider_name`.
|
|
1954
|
+
|
|
1955
|
+
For cloud providers: loads the API key from env var / keychain and
|
|
1956
|
+
raises `KeyNotConfiguredError` if missing.
|
|
1957
|
+
|
|
1958
|
+
For the local provider: skips the API-key path entirely. The URL is taken
|
|
1959
|
+
from `local_url` if provided, else from the keychain/env var, else
|
|
1960
|
+
`LocalProvider.DEFAULT_URL`.
|
|
1961
|
+
"""
|
|
1962
|
+
if provider_name not in PROVIDER_REGISTRY:
|
|
1963
|
+
raise UnknownProviderError(
|
|
1964
|
+
f"Provider '{provider_name}' is not yet wired up. "
|
|
1965
|
+
f"Currently supported: {', '.join(sorted(PROVIDER_REGISTRY))}."
|
|
1966
|
+
)
|
|
1967
|
+
|
|
1968
|
+
if provider_name == "local":
|
|
1969
|
+
url = local_url or load_local_url()
|
|
1970
|
+
return LocalProvider(base_url=url)
|
|
1971
|
+
|
|
1972
|
+
from cli_modelarium.security import load_key
|
|
1973
|
+
|
|
1974
|
+
api_key = load_key(provider_name)
|
|
1975
|
+
if not api_key:
|
|
1976
|
+
raise KeyNotConfiguredError(provider_name)
|
|
1977
|
+
|
|
1978
|
+
module_path, _, class_name = PROVIDER_REGISTRY[provider_name].partition(":")
|
|
1979
|
+
import importlib
|
|
1980
|
+
|
|
1981
|
+
module = importlib.import_module(module_path)
|
|
1982
|
+
provider_cls = getattr(module, class_name)
|
|
1983
|
+
return provider_cls(api_key=api_key)
|
|
1984
|
+
|
|
1985
|
+
|
|
1986
|
+
def _display_results(
|
|
1987
|
+
states: list[StreamState],
|
|
1988
|
+
judge_results: list[JudgeResult] | None = None,
|
|
1989
|
+
include_reasoning: bool = False,
|
|
1990
|
+
hallucination_mode: bool = False,
|
|
1991
|
+
hallucination_facts: list[str] | None = None,
|
|
1992
|
+
) -> None:
|
|
1993
|
+
"""Render the comparison results as a Rich table plus per-model output blocks.
|
|
1994
|
+
|
|
1995
|
+
`judge_results` is parallel to `states` when provided; it adds a Score
|
|
1996
|
+
column to the table and a Reasoning line under each output block when
|
|
1997
|
+
`include_reasoning=True`.
|
|
1998
|
+
|
|
1999
|
+
When `hallucination_mode=True`, the Score column is relabeled
|
|
2000
|
+
"Hallucination Risk" and each cell shows the worst-case panel risk
|
|
2001
|
+
plus the score (e.g. "Low (8)"). Color: Low=green, Medium=yellow,
|
|
2002
|
+
High=red.
|
|
2003
|
+
"""
|
|
2004
|
+
# Only surface the SP column when there are 2+ distinct non-empty
|
|
2005
|
+
# system prompts. The streaming legend has already printed the full
|
|
2006
|
+
# mapping; we just need the index here.
|
|
2007
|
+
from cli_modelarium.streaming import prompt_index_map
|
|
2008
|
+
|
|
2009
|
+
prompt_indices = prompt_index_map(states)
|
|
2010
|
+
show_sp_column = bool(prompt_indices)
|
|
2011
|
+
show_score_column = judge_results is not None
|
|
2012
|
+
|
|
2013
|
+
table = Table(
|
|
2014
|
+
title=f"Comparing {len(states)} completion{'s' if len(states) != 1 else ''}",
|
|
2015
|
+
border_style="dim",
|
|
2016
|
+
title_justify="left",
|
|
2017
|
+
)
|
|
2018
|
+
table.add_column("Model", style="bold")
|
|
2019
|
+
if show_sp_column:
|
|
2020
|
+
table.add_column("SP", style="magenta", justify="right")
|
|
2021
|
+
table.add_column("Temp", justify="right")
|
|
2022
|
+
table.add_column("TTFT", justify="right", style="dim")
|
|
2023
|
+
table.add_column("Latency", justify="right", style="dim")
|
|
2024
|
+
table.add_column("In", justify="right")
|
|
2025
|
+
table.add_column("Out", justify="right")
|
|
2026
|
+
table.add_column("Cost", justify="right")
|
|
2027
|
+
if show_score_column:
|
|
2028
|
+
score_header = "Hallucination Risk" if hallucination_mode else "Score"
|
|
2029
|
+
table.add_column(score_header, style="magenta", justify="right")
|
|
2030
|
+
table.add_column("Status")
|
|
2031
|
+
|
|
2032
|
+
total_cost = 0.0
|
|
2033
|
+
for i, s in enumerate(states):
|
|
2034
|
+
if s.error:
|
|
2035
|
+
status = "[red]error[/red]"
|
|
2036
|
+
cost_text = "[dim]-[/dim]"
|
|
2037
|
+
ttft_text = "[dim]-[/dim]"
|
|
2038
|
+
latency_text = "[dim]-[/dim]"
|
|
2039
|
+
in_text = "[dim]-[/dim]"
|
|
2040
|
+
out_text = "[dim]-[/dim]"
|
|
2041
|
+
else:
|
|
2042
|
+
status = "[green]ok[/green]"
|
|
2043
|
+
total_cost += s.cost_usd
|
|
2044
|
+
cost_text = "[dim]Free[/dim]" if is_local_model(s.model) else f"${s.cost_usd:.6f}"
|
|
2045
|
+
ttft_text = f"{s.ttft_ms / 1000:.2f}s" if s.ttft_ms is not None else "[dim]-[/dim]"
|
|
2046
|
+
latency_text = (
|
|
2047
|
+
f"{s.latency_ms / 1000:.2f}s" if s.latency_ms is not None else "[dim]-[/dim]"
|
|
2048
|
+
)
|
|
2049
|
+
in_text = str(s.input_tokens)
|
|
2050
|
+
out_text = str(s.output_tokens)
|
|
2051
|
+
|
|
2052
|
+
row: list[str] = [s.model]
|
|
2053
|
+
if show_sp_column:
|
|
2054
|
+
if s.system_prompt and s.system_prompt in prompt_indices:
|
|
2055
|
+
row.append(f"SP {prompt_indices[s.system_prompt]}")
|
|
2056
|
+
else:
|
|
2057
|
+
row.append("[dim]-[/dim]")
|
|
2058
|
+
row.extend(
|
|
2059
|
+
[
|
|
2060
|
+
f"{s.temperature:.1f}",
|
|
2061
|
+
ttft_text,
|
|
2062
|
+
latency_text,
|
|
2063
|
+
in_text,
|
|
2064
|
+
out_text,
|
|
2065
|
+
cost_text,
|
|
2066
|
+
]
|
|
2067
|
+
)
|
|
2068
|
+
if show_score_column:
|
|
2069
|
+
assert judge_results is not None
|
|
2070
|
+
if hallucination_mode:
|
|
2071
|
+
row.append(_risk_cell_for_compare(judge_results[i]))
|
|
2072
|
+
else:
|
|
2073
|
+
row.append(_score_cell_for_compare(judge_results[i]))
|
|
2074
|
+
row.append(status)
|
|
2075
|
+
table.add_row(*row)
|
|
2076
|
+
|
|
2077
|
+
console.print(table)
|
|
2078
|
+
console.print()
|
|
2079
|
+
|
|
2080
|
+
# Per-model output blocks. When multiple SPs are in play, identify which
|
|
2081
|
+
# one produced each block so the reader can cross-reference the legend.
|
|
2082
|
+
for i, s in enumerate(states):
|
|
2083
|
+
header = f"[bold cyan]>[/bold cyan] [bold]{s.model}[/bold] @ {s.temperature:.1f}"
|
|
2084
|
+
if show_sp_column and s.system_prompt and s.system_prompt in prompt_indices:
|
|
2085
|
+
header += f" [magenta]SP {prompt_indices[s.system_prompt]}[/magenta]"
|
|
2086
|
+
console.print(header)
|
|
2087
|
+
if s.error:
|
|
2088
|
+
console.print(f" [red]{s.error}[/red]")
|
|
2089
|
+
else:
|
|
2090
|
+
for line in s.text.splitlines() or [""]:
|
|
2091
|
+
console.print(f" {line}")
|
|
2092
|
+
# Optional judge reasoning lines.
|
|
2093
|
+
if include_reasoning and judge_results is not None:
|
|
2094
|
+
for j in judge_results[i].judges:
|
|
2095
|
+
score_str = j.score if j.score is not None else "?"
|
|
2096
|
+
if j.parse_error:
|
|
2097
|
+
console.print(
|
|
2098
|
+
f" [magenta dim]judge {j.model}: parse error - "
|
|
2099
|
+
f"{j.parse_error}[/magenta dim]"
|
|
2100
|
+
)
|
|
2101
|
+
else:
|
|
2102
|
+
console.print(
|
|
2103
|
+
f" [magenta dim]judge {j.model} ({score_str}/10): "
|
|
2104
|
+
f"{j.reasoning}[/magenta dim]"
|
|
2105
|
+
)
|
|
2106
|
+
console.print()
|
|
2107
|
+
|
|
2108
|
+
console.print(f"[dim]Total cost: ${total_cost:.6f}[/dim]")
|
|
2109
|
+
if judge_results is not None:
|
|
2110
|
+
j_cost = total_judge_cost(judge_results)
|
|
2111
|
+
j_calls = total_judge_calls(judge_results)
|
|
2112
|
+
console.print(
|
|
2113
|
+
f"[dim]Judge cost: ${j_cost:.6f} "
|
|
2114
|
+
f"({j_calls} judge call{'s' if j_calls != 1 else ''})[/dim]"
|
|
2115
|
+
)
|
|
2116
|
+
if hallucination_mode and hallucination_facts:
|
|
2117
|
+
console.print(
|
|
2118
|
+
f"[dim]Hallucination check: "
|
|
2119
|
+
f"{len(hallucination_facts)} reference fact"
|
|
2120
|
+
f"{'s' if len(hallucination_facts) != 1 else ''} provided[/dim]"
|
|
2121
|
+
)
|
|
2122
|
+
console.print(f"[dim]{pricing_freshness_note()}[/dim]")
|
|
2123
|
+
|
|
2124
|
+
|
|
2125
|
+
def _display_results_with_runs(
|
|
2126
|
+
states: list[StreamState],
|
|
2127
|
+
judge_results: list[JudgeResult] | None,
|
|
2128
|
+
runs: int,
|
|
2129
|
+
include_reasoning: bool = False,
|
|
2130
|
+
hallucination_mode: bool = False,
|
|
2131
|
+
hallucination_facts: list[str] | None = None,
|
|
2132
|
+
significance_results: list | None = None,
|
|
2133
|
+
stats_by_cell_cis: dict | None = None,
|
|
2134
|
+
mcnemar_results: list | None = None,
|
|
2135
|
+
) -> None:
|
|
2136
|
+
"""Render the runs > 1 path: one summary row per cell with RunStats.
|
|
2137
|
+
|
|
2138
|
+
Groups states by (model, temperature, system_prompt) cell, computes
|
|
2139
|
+
RunStats per cell, and prints a Rich table with statistical summary
|
|
2140
|
+
columns. Per-run outputs are shown below the table as a collapsed
|
|
2141
|
+
listing.
|
|
2142
|
+
|
|
2143
|
+
`judge_results` is parallel to `states` when provided. With mode-only
|
|
2144
|
+
judging, every state in a cell shares the same JudgeResult, so we pull
|
|
2145
|
+
the cell verdict from the first state.
|
|
2146
|
+
|
|
2147
|
+
When `hallucination_mode=True`, an additional "Hallucination Rate"
|
|
2148
|
+
summary is computed (fraction of runs flagged as High risk per cell).
|
|
2149
|
+
"""
|
|
2150
|
+
from cli_modelarium.run_statistics import compute_run_stats, group_states_by_cell
|
|
2151
|
+
|
|
2152
|
+
groups = group_states_by_cell(states)
|
|
2153
|
+
judge_by_state_id: dict[int, JudgeResult] = {}
|
|
2154
|
+
if judge_results is not None:
|
|
2155
|
+
for state, jr in zip(states, judge_results, strict=True):
|
|
2156
|
+
judge_by_state_id[id(state)] = jr
|
|
2157
|
+
|
|
2158
|
+
distinct_sps = {s.system_prompt for s in states if s.system_prompt}
|
|
2159
|
+
show_sp_column = len(distinct_sps) > 1
|
|
2160
|
+
show_hallucination_rate = hallucination_mode and judge_results is not None
|
|
2161
|
+
show_judge_column = judge_results is not None and not show_hallucination_rate
|
|
2162
|
+
|
|
2163
|
+
plural = "s" if len(groups) != 1 else ""
|
|
2164
|
+
title = f"Comparing {len(groups)} configuration{plural}, {runs} runs each"
|
|
2165
|
+
table = Table(title=title, border_style="dim", title_justify="left")
|
|
2166
|
+
table.add_column("Model", style="bold")
|
|
2167
|
+
if show_sp_column:
|
|
2168
|
+
table.add_column("SP", style="magenta", justify="right")
|
|
2169
|
+
table.add_column("Temp", justify="right")
|
|
2170
|
+
table.add_column("OK/Fail", justify="right")
|
|
2171
|
+
table.add_column("Latency mean ± stdev", justify="right", style="dim")
|
|
2172
|
+
table.add_column("CV", justify="right", style="dim")
|
|
2173
|
+
table.add_column("Tokens mean", justify="right")
|
|
2174
|
+
table.add_column("Cost total", justify="right")
|
|
2175
|
+
table.add_column("Diversity", justify="right")
|
|
2176
|
+
if show_hallucination_rate:
|
|
2177
|
+
table.add_column("Halluc. rate", justify="right", style="magenta")
|
|
2178
|
+
if show_judge_column:
|
|
2179
|
+
table.add_column("Score (mode)", justify="right", style="magenta")
|
|
2180
|
+
table.add_column("Mode", justify="left")
|
|
2181
|
+
|
|
2182
|
+
from cli_modelarium.streaming import prompt_index_map
|
|
2183
|
+
|
|
2184
|
+
prompt_indices = prompt_index_map(states)
|
|
2185
|
+
|
|
2186
|
+
grand_total_cost = 0.0
|
|
2187
|
+
cell_stats: list[tuple[tuple[str, float, str | None], list[StreamState], object]] = []
|
|
2188
|
+
for key, cell_states in groups.items():
|
|
2189
|
+
stats = compute_run_stats(cell_states)
|
|
2190
|
+
cell_stats.append((key, cell_states, stats))
|
|
2191
|
+
grand_total_cost += stats.cost_total_usd
|
|
2192
|
+
|
|
2193
|
+
model, temp, sp = key
|
|
2194
|
+
|
|
2195
|
+
# Hallucination rate: fraction of runs in this cell with risk_level "High".
|
|
2196
|
+
hallucination_rate_text = "[dim]-[/dim]"
|
|
2197
|
+
if show_hallucination_rate:
|
|
2198
|
+
high_count = 0
|
|
2199
|
+
judged = 0
|
|
2200
|
+
for s in cell_states:
|
|
2201
|
+
jr = judge_by_state_id.get(id(s))
|
|
2202
|
+
if jr is None or not jr.judges:
|
|
2203
|
+
continue
|
|
2204
|
+
judged += 1
|
|
2205
|
+
if jr.aggregated_risk_level == "High":
|
|
2206
|
+
high_count += 1
|
|
2207
|
+
if judged > 0:
|
|
2208
|
+
rate = high_count / judged
|
|
2209
|
+
color = (
|
|
2210
|
+
"red" if rate >= 0.5 else "yellow" if rate >= 0.2 else "green"
|
|
2211
|
+
)
|
|
2212
|
+
pct = f"{rate * 100:.0f}%"
|
|
2213
|
+
hallucination_rate_text = f"[{color}]{high_count}/{judged} ({pct})[/{color}]"
|
|
2214
|
+
|
|
2215
|
+
# Judge score (mode-only judging): pull the first non-empty JudgeResult.
|
|
2216
|
+
score_text = "[dim]-[/dim]"
|
|
2217
|
+
if show_judge_column:
|
|
2218
|
+
for s in cell_states:
|
|
2219
|
+
jr = judge_by_state_id.get(id(s))
|
|
2220
|
+
if jr is not None and jr.judges:
|
|
2221
|
+
score_text = _score_cell_for_compare(jr)
|
|
2222
|
+
break
|
|
2223
|
+
|
|
2224
|
+
if stats.latency_mean_ms is not None and stats.latency_stdev_ms is not None:
|
|
2225
|
+
latency_cell = f"{stats.latency_mean_ms:.0f} ± {stats.latency_stdev_ms:.0f} ms"
|
|
2226
|
+
elif stats.latency_mean_ms is not None:
|
|
2227
|
+
latency_cell = f"{stats.latency_mean_ms:.0f} ms"
|
|
2228
|
+
else:
|
|
2229
|
+
latency_cell = "[dim]-[/dim]"
|
|
2230
|
+
|
|
2231
|
+
cv_text = f"{stats.latency_cv:.3f}" if stats.latency_cv is not None else "[dim]-[/dim]"
|
|
2232
|
+
tokens_text = (
|
|
2233
|
+
f"{stats.output_tokens_mean:.0f}"
|
|
2234
|
+
if stats.output_tokens_mean is not None
|
|
2235
|
+
else "[dim]-[/dim]"
|
|
2236
|
+
)
|
|
2237
|
+
cost_text = (
|
|
2238
|
+
"[dim]Free[/dim]"
|
|
2239
|
+
if is_local_model(model)
|
|
2240
|
+
else f"${stats.cost_total_usd:.6f}"
|
|
2241
|
+
)
|
|
2242
|
+
diversity_text = f"{stats.output_diversity:.2f}"
|
|
2243
|
+
|
|
2244
|
+
if stats.mode_output is None:
|
|
2245
|
+
mode_text = "[dim]no mode (all unique)[/dim]"
|
|
2246
|
+
else:
|
|
2247
|
+
preview = stats.mode_output.replace("\n", " ").strip()
|
|
2248
|
+
if len(preview) > 50:
|
|
2249
|
+
preview = preview[:47] + "..."
|
|
2250
|
+
mode_text = f'"{preview}" ({stats.mode_count}x)'
|
|
2251
|
+
|
|
2252
|
+
row = [model]
|
|
2253
|
+
if show_sp_column:
|
|
2254
|
+
if sp and sp in prompt_indices:
|
|
2255
|
+
row.append(f"SP {prompt_indices[sp]}")
|
|
2256
|
+
else:
|
|
2257
|
+
row.append("[dim]-[/dim]")
|
|
2258
|
+
row.extend(
|
|
2259
|
+
[
|
|
2260
|
+
f"{temp:.1f}",
|
|
2261
|
+
f"{stats.n_succeeded}/{stats.n_failed}",
|
|
2262
|
+
latency_cell,
|
|
2263
|
+
cv_text,
|
|
2264
|
+
tokens_text,
|
|
2265
|
+
cost_text,
|
|
2266
|
+
diversity_text,
|
|
2267
|
+
]
|
|
2268
|
+
)
|
|
2269
|
+
if show_hallucination_rate:
|
|
2270
|
+
row.append(hallucination_rate_text)
|
|
2271
|
+
if show_judge_column:
|
|
2272
|
+
row.append(score_text)
|
|
2273
|
+
row.append(mode_text)
|
|
2274
|
+
table.add_row(*row)
|
|
2275
|
+
|
|
2276
|
+
console.print(table)
|
|
2277
|
+
console.print()
|
|
2278
|
+
|
|
2279
|
+
# Per-cell expanded view: list every run's output beneath its cell header.
|
|
2280
|
+
for key, cell_states, _stats in cell_stats:
|
|
2281
|
+
model, temp, sp = key
|
|
2282
|
+
header = f"[bold cyan]>[/bold cyan] [bold]{model}[/bold] @ {temp:.1f}"
|
|
2283
|
+
if show_sp_column and sp and sp in prompt_indices:
|
|
2284
|
+
header += f" [magenta]SP {prompt_indices[sp]}[/magenta]"
|
|
2285
|
+
console.print(header)
|
|
2286
|
+
for s in cell_states:
|
|
2287
|
+
tag = f" [dim]run {s.run_index + 1}/{runs}:[/dim]"
|
|
2288
|
+
if s.error:
|
|
2289
|
+
console.print(f"{tag} [red]{s.error}[/red]")
|
|
2290
|
+
else:
|
|
2291
|
+
lines = s.text.splitlines() or [""]
|
|
2292
|
+
console.print(f"{tag} {lines[0]}")
|
|
2293
|
+
for line in lines[1:]:
|
|
2294
|
+
console.print(f" {line}")
|
|
2295
|
+
if include_reasoning and judge_results is not None:
|
|
2296
|
+
for s in cell_states:
|
|
2297
|
+
jr = judge_by_state_id.get(id(s))
|
|
2298
|
+
if jr is None:
|
|
2299
|
+
continue
|
|
2300
|
+
for j in jr.judges:
|
|
2301
|
+
score_str = j.score if j.score is not None else "?"
|
|
2302
|
+
if j.parse_error:
|
|
2303
|
+
console.print(
|
|
2304
|
+
f" [magenta dim]judge {j.model}: parse error - "
|
|
2305
|
+
f"{j.parse_error}[/magenta dim]"
|
|
2306
|
+
)
|
|
2307
|
+
else:
|
|
2308
|
+
console.print(
|
|
2309
|
+
f" [magenta dim]judge {j.model} ({score_str}/10): "
|
|
2310
|
+
f"{j.reasoning}[/magenta dim]"
|
|
2311
|
+
)
|
|
2312
|
+
# Mode-only judging: one verdict per cell, no need to repeat.
|
|
2313
|
+
if runs > 1 and not hallucination_mode:
|
|
2314
|
+
break
|
|
2315
|
+
console.print()
|
|
2316
|
+
|
|
2317
|
+
console.print(f"[dim]Total cost across all runs: ${grand_total_cost:.6f}[/dim]")
|
|
2318
|
+
if judge_results is not None:
|
|
2319
|
+
j_cost = total_judge_cost(judge_results)
|
|
2320
|
+
j_calls = total_judge_calls(judge_results)
|
|
2321
|
+
console.print(
|
|
2322
|
+
f"[dim]Judge cost: ${j_cost:.6f} "
|
|
2323
|
+
f"({j_calls} judge call{'s' if j_calls != 1 else ''})[/dim]"
|
|
2324
|
+
)
|
|
2325
|
+
if hallucination_mode and hallucination_facts:
|
|
2326
|
+
console.print(
|
|
2327
|
+
f"[dim]Hallucination check: "
|
|
2328
|
+
f"{len(hallucination_facts)} reference fact"
|
|
2329
|
+
f"{'s' if len(hallucination_facts) != 1 else ''} provided[/dim]"
|
|
2330
|
+
)
|
|
2331
|
+
console.print(
|
|
2332
|
+
"[dim]Coefficient of variation (CV) < 0.05 indicates stable model behavior.[/dim]"
|
|
2333
|
+
)
|
|
2334
|
+
console.print(f"[dim]{pricing_freshness_note()}[/dim]")
|
|
2335
|
+
|
|
2336
|
+
if stats_by_cell_cis:
|
|
2337
|
+
_display_confidence_intervals(stats_by_cell_cis)
|
|
2338
|
+
|
|
2339
|
+
if significance_results:
|
|
2340
|
+
_display_significance(significance_results)
|
|
2341
|
+
|
|
2342
|
+
if mcnemar_results:
|
|
2343
|
+
_display_mcnemar(mcnemar_results)
|
|
2344
|
+
|
|
2345
|
+
|
|
2346
|
+
def _display_confidence_intervals(stats_by_cell_cis: dict) -> None:
|
|
2347
|
+
"""Render bootstrap CIs on per-model means below the runs table."""
|
|
2348
|
+
if not stats_by_cell_cis:
|
|
2349
|
+
return
|
|
2350
|
+
has_any = any(metrics for metrics in stats_by_cell_cis.values())
|
|
2351
|
+
if not has_any:
|
|
2352
|
+
return
|
|
2353
|
+
|
|
2354
|
+
console.print()
|
|
2355
|
+
console.print("[bold]Bootstrap Confidence Intervals[/bold]")
|
|
2356
|
+
for model, metrics in stats_by_cell_cis.items():
|
|
2357
|
+
if not metrics:
|
|
2358
|
+
continue
|
|
2359
|
+
parts: list[str] = []
|
|
2360
|
+
for metric_name in ("latency_ms", "score", "output_tokens", "cost_usd"):
|
|
2361
|
+
ci = metrics.get(metric_name)
|
|
2362
|
+
if ci is None:
|
|
2363
|
+
continue
|
|
2364
|
+
label = {
|
|
2365
|
+
"latency_ms": "latency",
|
|
2366
|
+
"score": "score",
|
|
2367
|
+
"output_tokens": "tokens",
|
|
2368
|
+
"cost_usd": "cost",
|
|
2369
|
+
}[metric_name]
|
|
2370
|
+
level_pct = int(round(ci["ci_level"] * 100))
|
|
2371
|
+
parts.append(
|
|
2372
|
+
f"{label} [{level_pct}% CI: {ci['ci_low']:.3f}, {ci['ci_high']:.3f}]"
|
|
2373
|
+
)
|
|
2374
|
+
if parts:
|
|
2375
|
+
console.print(f" {model}: " + " | ".join(parts))
|
|
2376
|
+
|
|
2377
|
+
|
|
2378
|
+
def _display_mcnemar(mcnemar_results: list) -> None:
|
|
2379
|
+
"""Render McNemar test results for paired binary outcomes."""
|
|
2380
|
+
if not mcnemar_results:
|
|
2381
|
+
return
|
|
2382
|
+
|
|
2383
|
+
console.print()
|
|
2384
|
+
console.print("[bold]Binary Outcome Significance (McNemar)[/bold]")
|
|
2385
|
+
first = mcnemar_results[0]
|
|
2386
|
+
console.print(
|
|
2387
|
+
f"[dim]Metric: hallucination pass/fail | Correction: "
|
|
2388
|
+
f"{first.correction_method} | Threshold: p < {first.threshold}[/dim]"
|
|
2389
|
+
)
|
|
2390
|
+
|
|
2391
|
+
for r in mcnemar_results:
|
|
2392
|
+
if r.n_discordant == 0:
|
|
2393
|
+
console.print(
|
|
2394
|
+
f" {r.model_a} ({r.a_pass_rate:.0%} pass) vs "
|
|
2395
|
+
f"{r.model_b} ({r.b_pass_rate:.0%} pass): "
|
|
2396
|
+
f"no discordant runs (test undefined)"
|
|
2397
|
+
)
|
|
2398
|
+
continue
|
|
2399
|
+
sig_marker = "*" if r.significant_at_threshold else ""
|
|
2400
|
+
p_display = (
|
|
2401
|
+
r.p_value_corrected if r.p_value_corrected is not None else r.p_value
|
|
2402
|
+
)
|
|
2403
|
+
method_label = {
|
|
2404
|
+
"exact_binomial": "exact",
|
|
2405
|
+
"edwards_chi2": "Edwards",
|
|
2406
|
+
}.get(r.method, r.method)
|
|
2407
|
+
console.print(
|
|
2408
|
+
f" {r.model_a} ({r.a_pass_rate:.0%} pass) vs "
|
|
2409
|
+
f"{r.model_b} ({r.b_pass_rate:.0%} pass): "
|
|
2410
|
+
f"p={p_display:.4f}{sig_marker} "
|
|
2411
|
+
f"({method_label}, discordant={r.n_discordant})"
|
|
2412
|
+
)
|
|
2413
|
+
|
|
2414
|
+
|
|
2415
|
+
def _display_significance(significance_results: list) -> None:
|
|
2416
|
+
"""Render pairwise statistical significance results below the runs table.
|
|
2417
|
+
|
|
2418
|
+
Display strategy depends on the number of models:
|
|
2419
|
+
* 2 models: single-line summary
|
|
2420
|
+
* 3-5 models: matrix table
|
|
2421
|
+
* 6+ models: top-K significant pairs (full matrix in JSON)
|
|
2422
|
+
"""
|
|
2423
|
+
if not significance_results:
|
|
2424
|
+
return
|
|
2425
|
+
|
|
2426
|
+
models = sorted(
|
|
2427
|
+
{r.model_a for r in significance_results}
|
|
2428
|
+
| {r.model_b for r in significance_results}
|
|
2429
|
+
)
|
|
2430
|
+
n_models = len(models)
|
|
2431
|
+
first = significance_results[0]
|
|
2432
|
+
|
|
2433
|
+
console.print()
|
|
2434
|
+
console.print("[bold]Statistical Significance Tests[/bold]")
|
|
2435
|
+
console.print(
|
|
2436
|
+
f"[dim]Metric: {first.metric} | Test: {first.test_used} | "
|
|
2437
|
+
f"Correction: {first.correction_method} | Threshold: p < {first.threshold}[/dim]"
|
|
2438
|
+
)
|
|
2439
|
+
|
|
2440
|
+
if n_models == 2:
|
|
2441
|
+
r = significance_results[0]
|
|
2442
|
+
if r.p_value is None:
|
|
2443
|
+
console.print(
|
|
2444
|
+
f" {r.model_a} vs {r.model_b}: {r.test_used} (no p-value)"
|
|
2445
|
+
)
|
|
2446
|
+
else:
|
|
2447
|
+
sig_marker = "*" if r.significant_at_threshold else ""
|
|
2448
|
+
p_display = (
|
|
2449
|
+
r.p_value_corrected if r.p_value_corrected is not None else r.p_value
|
|
2450
|
+
)
|
|
2451
|
+
d_text = (
|
|
2452
|
+
f", d={r.effect_size:.3f} ({r.effect_size_interpretation})"
|
|
2453
|
+
if r.effect_size is not None
|
|
2454
|
+
else ""
|
|
2455
|
+
)
|
|
2456
|
+
console.print(
|
|
2457
|
+
f" {r.model_a} (avg {r.mean_a:.3f}) vs "
|
|
2458
|
+
f"{r.model_b} (avg {r.mean_b:.3f}): "
|
|
2459
|
+
f"p={p_display:.4f}{sig_marker}{d_text}"
|
|
2460
|
+
)
|
|
2461
|
+
return
|
|
2462
|
+
|
|
2463
|
+
if n_models <= 5:
|
|
2464
|
+
table = Table(title="Pairwise p-values (corrected)", border_style="dim")
|
|
2465
|
+
table.add_column("Model", style="cyan")
|
|
2466
|
+
for m in models:
|
|
2467
|
+
table.add_column(m, justify="right")
|
|
2468
|
+
|
|
2469
|
+
result_map: dict[tuple[str, str], object] = {}
|
|
2470
|
+
for r in significance_results:
|
|
2471
|
+
result_map[(r.model_a, r.model_b)] = r
|
|
2472
|
+
result_map[(r.model_b, r.model_a)] = r
|
|
2473
|
+
|
|
2474
|
+
for m_a in models:
|
|
2475
|
+
row = [m_a]
|
|
2476
|
+
for m_b in models:
|
|
2477
|
+
if m_a == m_b:
|
|
2478
|
+
row.append("-")
|
|
2479
|
+
continue
|
|
2480
|
+
r = result_map.get((m_a, m_b))
|
|
2481
|
+
if r is None or r.p_value is None: # type: ignore[union-attr]
|
|
2482
|
+
row.append("-")
|
|
2483
|
+
else:
|
|
2484
|
+
p = (
|
|
2485
|
+
r.p_value_corrected # type: ignore[union-attr]
|
|
2486
|
+
if r.p_value_corrected is not None # type: ignore[union-attr]
|
|
2487
|
+
else r.p_value # type: ignore[union-attr]
|
|
2488
|
+
)
|
|
2489
|
+
marker = "*" if r.significant_at_threshold else "" # type: ignore[union-attr]
|
|
2490
|
+
row.append(f"{p:.4f}{marker}")
|
|
2491
|
+
table.add_row(*row)
|
|
2492
|
+
|
|
2493
|
+
console.print(table)
|
|
2494
|
+
console.print("[dim]* = significant after correction[/dim]")
|
|
2495
|
+
return
|
|
2496
|
+
|
|
2497
|
+
# 6+ models: top-K significant
|
|
2498
|
+
significant = [r for r in significance_results if r.significant_at_threshold]
|
|
2499
|
+
significant.sort(key=lambda x: x.p_value_corrected or 1.0)
|
|
2500
|
+
top_k = significant[:5]
|
|
2501
|
+
|
|
2502
|
+
if top_k:
|
|
2503
|
+
console.print(
|
|
2504
|
+
f"[bold]Top significant pairs (of {len(significant)} total):[/bold]"
|
|
2505
|
+
)
|
|
2506
|
+
for i, r in enumerate(top_k, 1):
|
|
2507
|
+
p = r.p_value_corrected if r.p_value_corrected is not None else r.p_value
|
|
2508
|
+
d_text = (
|
|
2509
|
+
f", d={r.effect_size:.3f} ({r.effect_size_interpretation})"
|
|
2510
|
+
if r.effect_size is not None
|
|
2511
|
+
else ""
|
|
2512
|
+
)
|
|
2513
|
+
console.print(
|
|
2514
|
+
f" {i}. {r.model_a} vs {r.model_b}: p={p:.4f}{d_text}"
|
|
2515
|
+
)
|
|
2516
|
+
else:
|
|
2517
|
+
console.print("[dim]No statistically significant pairs found.[/dim]")
|
|
2518
|
+
|
|
2519
|
+
console.print("[dim]Full matrix available in JSON output.[/dim]")
|
|
2520
|
+
|
|
2521
|
+
|
|
2522
|
+
def _score_cell_for_compare(jr: JudgeResult) -> str:
|
|
2523
|
+
"""Render the Score column cell for the compare command's results table."""
|
|
2524
|
+
if not jr.judges:
|
|
2525
|
+
if jr.skipped_models:
|
|
2526
|
+
return "[dim]-[/dim]"
|
|
2527
|
+
return "[dim]-[/dim]"
|
|
2528
|
+
successful = [j for j in jr.judges if j.score is not None]
|
|
2529
|
+
if not successful:
|
|
2530
|
+
return "[red]N/A[/red]"
|
|
2531
|
+
if len(jr.judges) == 1 and successful:
|
|
2532
|
+
# Single judge: just the score.
|
|
2533
|
+
return str(successful[0].score)
|
|
2534
|
+
# Panel: average + count.
|
|
2535
|
+
if jr.average_score is not None:
|
|
2536
|
+
return f"{jr.average_score:.1f} ({len(successful)})"
|
|
2537
|
+
return "[red]N/A[/red]"
|
|
2538
|
+
|
|
2539
|
+
|
|
2540
|
+
_RISK_COLOR = {"Low": "green", "Medium": "yellow", "High": "red"}
|
|
2541
|
+
|
|
2542
|
+
|
|
2543
|
+
def _risk_cell_for_compare(jr: JudgeResult) -> str:
|
|
2544
|
+
"""Render the Hallucination Risk cell for one row.
|
|
2545
|
+
|
|
2546
|
+
Single judge: "[Low] (8)" colorized by risk level.
|
|
2547
|
+
Panel: worst-case risk + score range, e.g. "[High] (3-7)".
|
|
2548
|
+
No data: dim dash.
|
|
2549
|
+
"""
|
|
2550
|
+
if not jr.judges:
|
|
2551
|
+
return "[dim]-[/dim]"
|
|
2552
|
+
successful = [j for j in jr.judges if j.score is not None and j.risk_level]
|
|
2553
|
+
if not successful:
|
|
2554
|
+
return "[red]N/A[/red]"
|
|
2555
|
+
|
|
2556
|
+
risk = jr.aggregated_risk_level or "?"
|
|
2557
|
+
color = _RISK_COLOR.get(risk, "magenta")
|
|
2558
|
+
|
|
2559
|
+
if len(successful) == 1:
|
|
2560
|
+
s = successful[0]
|
|
2561
|
+
return f"[{color}]{s.risk_level}[/{color}] ({s.score})"
|
|
2562
|
+
|
|
2563
|
+
scores = sorted(j.score for j in successful if j.score is not None)
|
|
2564
|
+
if scores[0] == scores[-1]:
|
|
2565
|
+
score_text = str(scores[0])
|
|
2566
|
+
else:
|
|
2567
|
+
score_text = f"{scores[0]}-{scores[-1]}"
|
|
2568
|
+
return f"[{color}]{risk}[/{color}] ({score_text}, n={len(successful)})"
|
|
2569
|
+
|
|
2570
|
+
|
|
2571
|
+
def _print_error(message: str) -> None:
|
|
2572
|
+
"""Print an error inside a red-bordered panel."""
|
|
2573
|
+
console.print(Panel(redact_secrets(message), title="Error", border_style="red"))
|
|
2574
|
+
|
|
2575
|
+
|
|
2576
|
+
if __name__ == "__main__":
|
|
2577
|
+
main()
|