auditkit 1.0.0__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.
- auditkit/README.md +99 -0
- auditkit/__init__.py +177 -0
- auditkit/__main__.py +3 -0
- auditkit/_bootstrap.py +77 -0
- auditkit/_identity_guard.py +99 -0
- auditkit/adapter.py +264 -0
- auditkit/annotator.py +339 -0
- auditkit/api.py +502 -0
- auditkit/assets/auditkit_logo.png +0 -0
- auditkit/cache.py +47 -0
- auditkit/cli.py +417 -0
- auditkit/comparison.py +563 -0
- auditkit/diff.py +265 -0
- auditkit/errors.py +54 -0
- auditkit/evaluator.py +20 -0
- auditkit/experiment.py +145 -0
- auditkit/hf_publish.py +262 -0
- auditkit/lmeval_engine.py +550 -0
- auditkit/loaders.py +121 -0
- auditkit/logs.py +18 -0
- auditkit/metric.py +199 -0
- auditkit/metrics/README.md +15 -0
- auditkit/metrics/__init__.py +0 -0
- auditkit/metrics/code.py +222 -0
- auditkit/metrics/embedding.py +131 -0
- auditkit/metrics/encoder_judge.py +423 -0
- auditkit/metrics/generation.py +331 -0
- auditkit/metrics/guard.py +412 -0
- auditkit/metrics/hallucination.py +45 -0
- auditkit/metrics/judge.py +547 -0
- auditkit/metrics/pairwise.py +153 -0
- auditkit/metrics/perf.py +53 -0
- auditkit/metrics/rag.py +149 -0
- auditkit/metrics/security.py +64 -0
- auditkit/metrics/toxicity.py +238 -0
- auditkit/model/README.md +16 -0
- auditkit/model/__init__.py +485 -0
- auditkit/model/anthropic.py +94 -0
- auditkit/model/api_gen.py +133 -0
- auditkit/model/groq_gen.py +121 -0
- auditkit/model/hf_gen.py +385 -0
- auditkit/model/lexsi.py +155 -0
- auditkit/model/litellm_gen.py +65 -0
- auditkit/model/openai.py +90 -0
- auditkit/model/openrouter_gen.py +152 -0
- auditkit/model/vllm_gen.py +316 -0
- auditkit/model_compare.py +655 -0
- auditkit/redteam/README.md +9 -0
- auditkit/redteam/__init__.py +26 -0
- auditkit/redteam/detector.py +37 -0
- auditkit/redteam/detectors/README.md +5 -0
- auditkit/redteam/detectors/builtin.py +126 -0
- auditkit/redteam/probe.py +39 -0
- auditkit/redteam/probes/README.md +5 -0
- auditkit/redteam/probes/builtin.py +85 -0
- auditkit/redteam/runner.py +206 -0
- auditkit/registry.py +65 -0
- auditkit/report.py +278 -0
- auditkit/report_format.py +52 -0
- auditkit/router.py +54 -0
- auditkit/runner.py +575 -0
- auditkit/runspec.py +159 -0
- auditkit/sample.py +40 -0
- auditkit/scenario.py +88 -0
- auditkit/scenarios/README.md +10 -0
- auditkit/scenarios/__init__.py +4 -0
- auditkit/scenarios/arc.py +33 -0
- auditkit/scenarios/gsm8k.py +32 -0
- auditkit/scenarios/hellaswag.py +33 -0
- auditkit/scenarios/humaneval.py +32 -0
- auditkit/scenarios/mmlu.py +34 -0
- auditkit/scenarios/truthfulqa.py +33 -0
- auditkit/score.py +165 -0
- auditkit/scorers.py +117 -0
- auditkit/scoring.py +79 -0
- auditkit/types.py +69 -0
- auditkit-1.0.0.dist-info/METADATA +396 -0
- auditkit-1.0.0.dist-info/RECORD +81 -0
- auditkit-1.0.0.dist-info/WHEEL +4 -0
- auditkit-1.0.0.dist-info/entry_points.txt +2 -0
- auditkit-1.0.0.dist-info/licenses/LICENSE.md +92 -0
|
@@ -0,0 +1,550 @@
|
|
|
1
|
+
"""The benchmark engine: run lm-evaluation-harness with its *full* task machinery.
|
|
2
|
+
|
|
3
|
+
Unlike the native spine (which owns prompting and scoring), lm-eval is a whole
|
|
4
|
+
pipeline — task + prompt template + output_type + filters + metrics + its own
|
|
5
|
+
model backends. So it is wrapped here as a *technique engine*, not a
|
|
6
|
+
:class:`~auditkit.model.Model`: :func:`run_benchmark` maps our model spec onto
|
|
7
|
+
lm-eval's backend, calls ``lm_eval.simple_evaluate(..., log_samples=True)``, and
|
|
8
|
+
maps the aggregate results and per-doc logged samples back into our uniform
|
|
9
|
+
:class:`~auditkit.report.RunResult` (headline + per-metric ``Stat`` + one
|
|
10
|
+
:class:`~auditkit.report.Prediction` per sample, i.e. the answer browser).
|
|
11
|
+
|
|
12
|
+
lm-eval is an optional extra: ``pip install auditkit[lmeval]``. The import is
|
|
13
|
+
lazy, so this module imports fine without it.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import contextlib
|
|
19
|
+
import hashlib
|
|
20
|
+
import json
|
|
21
|
+
import os
|
|
22
|
+
from typing import Any, Optional, Union
|
|
23
|
+
|
|
24
|
+
from .cache import DiskCache
|
|
25
|
+
from .errors import AuditKitError, CapabilityError, ExtraNotInstalled
|
|
26
|
+
from .report import Prediction, RunResult
|
|
27
|
+
from .score import Stat
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@contextlib.contextmanager
|
|
31
|
+
def _hf_token_env(token: Optional[str]):
|
|
32
|
+
"""Temporarily set ``HF_TOKEN`` for the enclosed call, then restore it.
|
|
33
|
+
|
|
34
|
+
Passing ``token`` in lm-eval's ``model_args`` covers gated *model* downloads,
|
|
35
|
+
but gated *datasets* (some task data) resolve through ``datasets``, which
|
|
36
|
+
reads ``HF_TOKEN`` from the environment. This forwards an explicitly-passed
|
|
37
|
+
token to that path for the duration of the run only — no ambient env is read,
|
|
38
|
+
and the previous value is always restored.
|
|
39
|
+
"""
|
|
40
|
+
if not token:
|
|
41
|
+
yield
|
|
42
|
+
return
|
|
43
|
+
prev = os.environ.get("HF_TOKEN")
|
|
44
|
+
os.environ["HF_TOKEN"] = token
|
|
45
|
+
try:
|
|
46
|
+
yield
|
|
47
|
+
finally:
|
|
48
|
+
if prev is None:
|
|
49
|
+
os.environ.pop("HF_TOKEN", None)
|
|
50
|
+
else:
|
|
51
|
+
os.environ["HF_TOKEN"] = prev
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@contextlib.contextmanager
|
|
55
|
+
def _temp_env(name: str, value: Optional[str]):
|
|
56
|
+
"""Temporarily set ``os.environ[name] = value`` for the enclosed call, then
|
|
57
|
+
restore the previous value (or unset). No-op when ``value`` is falsy. Used to
|
|
58
|
+
mirror a Groq key into ``OPENAI_API_KEY`` for the openai-chat-completions
|
|
59
|
+
backend without touching the ambient environment beyond the run."""
|
|
60
|
+
if not value:
|
|
61
|
+
yield
|
|
62
|
+
return
|
|
63
|
+
prev = os.environ.get(name)
|
|
64
|
+
os.environ[name] = value
|
|
65
|
+
try:
|
|
66
|
+
yield
|
|
67
|
+
finally:
|
|
68
|
+
if prev is None:
|
|
69
|
+
os.environ.pop(name, None)
|
|
70
|
+
else:
|
|
71
|
+
os.environ[name] = prev
|
|
72
|
+
|
|
73
|
+
# Our model-spec prefix → (lm-eval backend name, the arg key that holds the model name).
|
|
74
|
+
_BACKEND_MAP: dict[str, tuple[str, str]] = {
|
|
75
|
+
"hf": ("hf", "pretrained"),
|
|
76
|
+
"vllm": ("vllm", "pretrained"),
|
|
77
|
+
"openai": ("openai-chat-completions", "model"),
|
|
78
|
+
"anthropic": ("anthropic-chat", "model"),
|
|
79
|
+
"groq": ("openai-chat-completions", "model"), # OpenAI-compatible chat endpoint (see below)
|
|
80
|
+
"openrouter": ("openai-chat-completions", "model"), # same shape as groq, see below
|
|
81
|
+
"api": ("local-completions", "model"), # OpenAI-compatible server (logprob-capable)
|
|
82
|
+
"lexsi": ("local-completions", "model"), # Lexsi gateway is OpenAI-compatible
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
# Groq speaks the OpenAI *chat* API (no /completions, no logprobs), so it rides
|
|
86
|
+
# the openai-chat-completions backend pointed at this base URL. It authenticates
|
|
87
|
+
# via the OPENAI_API_KEY env (that backend's api_key is a read-only property, not
|
|
88
|
+
# a model_arg), so run_benchmark() mirrors GROQ_API_KEY/api_key into it for the
|
|
89
|
+
# duration of the run only. Being openai-chat-completions, it's already in
|
|
90
|
+
# _CHAT_ONLY, so loglikelihood/MC tasks are rejected up front.
|
|
91
|
+
_GROQ_BASE_URL = "https://api.groq.com/openai/v1/chat/completions"
|
|
92
|
+
|
|
93
|
+
# OpenRouter is the same shape as Groq -- an OpenAI-compatible chat endpoint,
|
|
94
|
+
# authenticated the same way (mirrored into OPENAI_API_KEY for the run only).
|
|
95
|
+
_OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1/chat/completions"
|
|
96
|
+
|
|
97
|
+
# Backends that cannot return logprobs, so loglikelihood/MC tasks can't run on them.
|
|
98
|
+
_CHAT_ONLY = {"openai-chat-completions", "anthropic-chat"}
|
|
99
|
+
|
|
100
|
+
# lm-eval output types that are scored by loglikelihood (need a logprob-capable backend).
|
|
101
|
+
_LOGLIKELIHOOD_TYPES = {"loglikelihood", "multiple_choice", "loglikelihood_rolling"}
|
|
102
|
+
|
|
103
|
+
# Passed straight through to lm-eval's model_args when present in opts.
|
|
104
|
+
_PASSTHROUGH_ARGS = (
|
|
105
|
+
"base_url", "api_key", "dtype", "device", "trust_remote_code",
|
|
106
|
+
"tokenizer", "revision", "tensor_parallel_size", "max_length", "peft",
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
# Per-sample metric keys lm-eval commonly logs, in preference order for "correct".
|
|
110
|
+
_KNOWN_SAMPLE_METRICS = ("acc", "acc_norm", "exact_match", "em", "f1", "mc1", "mc2")
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def map_model_spec(model: Any, **opts: Any) -> tuple[str, dict[str, Any]]:
|
|
114
|
+
"""Map an AuditKIT model spec to an lm-eval ``(backend, model_args)`` pair.
|
|
115
|
+
|
|
116
|
+
Accepts a string like ``"hf:gpt2"``, ``"vllm:meta-llama/…"``,
|
|
117
|
+
``"openai:gpt-4o"``, ``"groq:llama-3.3-70b-versatile"``, or ``"api:my-model"``
|
|
118
|
+
(with ``base_url=``). A bare name is treated as ``hf``. A callable/:class:`Model`
|
|
119
|
+
is rejected — lm-eval builds its own backend, so those must use the native
|
|
120
|
+
engine.
|
|
121
|
+
|
|
122
|
+
``groq:`` routes through the openai-chat-completions backend with Groq's base
|
|
123
|
+
URL filled in automatically; it's chat-only, so generation-scored tasks only.
|
|
124
|
+
"""
|
|
125
|
+
if not isinstance(model, str):
|
|
126
|
+
raise AuditKitError(
|
|
127
|
+
f"the benchmark engine needs a string model spec (e.g. 'hf:gpt2', "
|
|
128
|
+
f"'vllm:...', 'api:name' with base_url); got {type(model).__name__}. "
|
|
129
|
+
f"lm-eval constructs its own backend — use engine='native' for a "
|
|
130
|
+
f"callable or Model instance."
|
|
131
|
+
)
|
|
132
|
+
if ":" in model:
|
|
133
|
+
prefix, name = model.split(":", 1)
|
|
134
|
+
else:
|
|
135
|
+
prefix, name = "hf", model
|
|
136
|
+
if prefix not in _BACKEND_MAP:
|
|
137
|
+
raise AuditKitError(
|
|
138
|
+
f"model prefix {prefix!r} is not supported by the benchmark engine; "
|
|
139
|
+
f"known: {sorted(_BACKEND_MAP)}"
|
|
140
|
+
)
|
|
141
|
+
backend, name_key = _BACKEND_MAP[prefix]
|
|
142
|
+
args: dict[str, Any] = {name_key: name}
|
|
143
|
+
for k in _PASSTHROUGH_ARGS:
|
|
144
|
+
if opts.get(k) is not None:
|
|
145
|
+
args[k] = opts[k]
|
|
146
|
+
# HF gated models (Llama, Gemma, ...): accept `hf_token` or `token` and pass
|
|
147
|
+
# it to lm-eval as the `token` model_arg (hf/vllm read it for from_pretrained).
|
|
148
|
+
# run_benchmark() also mirrors it into HF_TOKEN so gated *datasets* resolve.
|
|
149
|
+
hf_token = opts.get("hf_token") or opts.get("token")
|
|
150
|
+
if hf_token is not None:
|
|
151
|
+
args["token"] = hf_token
|
|
152
|
+
if prefix == "groq":
|
|
153
|
+
# Fill in Groq's endpoint (a caller-supplied base_url still wins). The
|
|
154
|
+
# key is authenticated via OPENAI_API_KEY env in run_benchmark, never a
|
|
155
|
+
# model_arg -- so drop any api_key that flowed through the passthrough.
|
|
156
|
+
args.pop("api_key", None)
|
|
157
|
+
args.setdefault("base_url", _GROQ_BASE_URL)
|
|
158
|
+
if prefix == "openrouter":
|
|
159
|
+
# Same reasoning as groq above.
|
|
160
|
+
args.pop("api_key", None)
|
|
161
|
+
args.setdefault("base_url", _OPENROUTER_BASE_URL)
|
|
162
|
+
if backend == "local-completions" and "base_url" not in args:
|
|
163
|
+
raise AuditKitError(
|
|
164
|
+
f"'{prefix}:' needs base_url=<OpenAI-compatible endpoint> for the "
|
|
165
|
+
f"benchmark engine (e.g. base_url='https://…/v1/completions')."
|
|
166
|
+
)
|
|
167
|
+
return backend, args
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def _task_output_types(tasks: list[str]) -> dict[str, str]:
|
|
171
|
+
"""Best-effort map of task name → lm-eval output_type (empty if unknowable)."""
|
|
172
|
+
from lm_eval.tasks import TaskManager, get_task_dict
|
|
173
|
+
|
|
174
|
+
td = get_task_dict(tasks, TaskManager())
|
|
175
|
+
types: dict[str, str] = {}
|
|
176
|
+
|
|
177
|
+
def walk(d: dict) -> None:
|
|
178
|
+
for k, v in d.items():
|
|
179
|
+
if isinstance(v, dict):
|
|
180
|
+
walk(v)
|
|
181
|
+
continue
|
|
182
|
+
ot = getattr(v, "OUTPUT_TYPE", None)
|
|
183
|
+
if ot is None:
|
|
184
|
+
cfg = getattr(v, "config", None)
|
|
185
|
+
ot = getattr(cfg, "output_type", None)
|
|
186
|
+
if ot:
|
|
187
|
+
types[k] = ot
|
|
188
|
+
|
|
189
|
+
walk(td)
|
|
190
|
+
return types
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _assert_chat_compatible(tasks: list[str], model: str) -> None:
|
|
194
|
+
"""Raise if a chat-only model is asked to run a loglikelihood/MC task."""
|
|
195
|
+
try:
|
|
196
|
+
types = _task_output_types(tasks)
|
|
197
|
+
except Exception:
|
|
198
|
+
return # can't introspect; the post-call error catch is the safety net
|
|
199
|
+
bad = sorted(t for t, ot in types.items() if ot in _LOGLIKELIHOOD_TYPES)
|
|
200
|
+
if bad:
|
|
201
|
+
raise CapabilityError(
|
|
202
|
+
f"task(s) {bad} score by loglikelihood, but {model!r} is chat-only "
|
|
203
|
+
f"(no logprobs). Use a logprob-capable model (hf:/vllm:/api: with "
|
|
204
|
+
f"base_url), or a generation-scored task variant."
|
|
205
|
+
)
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def _looks_like_loglikelihood_error(exc: Exception) -> bool:
|
|
209
|
+
s = f"{type(exc).__name__}: {exc}".lower()
|
|
210
|
+
return "loglikelihood" in s or "not implement" in s
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
# lm-eval run-level knobs we forward explicitly (everything else goes via
|
|
214
|
+
# lmeval_kwargs). Value None means "don't pass it — use lm-eval's default".
|
|
215
|
+
_RUN_KNOBS = ("apply_chat_template", "system_instruction", "gen_kwargs",
|
|
216
|
+
"fewshot_as_multiturn", "use_cache", "cache_requests")
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
def run_benchmark(
|
|
220
|
+
tasks: Union[str, list[str]],
|
|
221
|
+
model: Any,
|
|
222
|
+
*,
|
|
223
|
+
config: Any = None,
|
|
224
|
+
run_name: str = "",
|
|
225
|
+
model_args: Optional[dict[str, Any]] = None,
|
|
226
|
+
apply_chat_template: Optional[bool] = None,
|
|
227
|
+
system_instruction: Optional[str] = None,
|
|
228
|
+
gen_kwargs: Any = None,
|
|
229
|
+
fewshot_as_multiturn: Optional[bool] = None,
|
|
230
|
+
lmeval_kwargs: Optional[dict[str, Any]] = None,
|
|
231
|
+
**opts: Any,
|
|
232
|
+
) -> RunResult:
|
|
233
|
+
"""Run one or more lm-eval tasks and return a unified :class:`RunResult`.
|
|
234
|
+
|
|
235
|
+
``tasks`` is a task name, comma-separated string, or list of names. ``model``
|
|
236
|
+
is a string spec mapped by :func:`map_model_spec`. ``config`` is a
|
|
237
|
+
:class:`~auditkit.runspec.RunConfig` (``num_fewshot``, ``limit``, ``seed``,
|
|
238
|
+
``batch_size`` are honored).
|
|
239
|
+
|
|
240
|
+
Extra ``opts`` (``base_url``, ``dtype``, ``device``, ``hf_token``, …) flow
|
|
241
|
+
into lm-eval's ``model_args`` (the model backend). lm-eval *run* knobs are
|
|
242
|
+
exposed directly — ``apply_chat_template`` (set True for instruct/chat/
|
|
243
|
+
fine-tuned models), ``system_instruction``, ``gen_kwargs`` (generation params
|
|
244
|
+
for generative tasks, e.g. ``"temperature=0,max_gen_toks=256"``),
|
|
245
|
+
``fewshot_as_multiturn`` — and anything else lm-eval's ``simple_evaluate``
|
|
246
|
+
accepts can be passed via ``lmeval_kwargs={...}`` for full parity.
|
|
247
|
+
"""
|
|
248
|
+
from .runspec import RunConfig
|
|
249
|
+
|
|
250
|
+
if isinstance(tasks, str):
|
|
251
|
+
tasks = [t.strip() for t in tasks.split(",") if t.strip()]
|
|
252
|
+
if not tasks:
|
|
253
|
+
raise AuditKitError("benchmark engine needs at least one task name")
|
|
254
|
+
cfg = config or RunConfig()
|
|
255
|
+
|
|
256
|
+
backend, resolved_model_args = map_model_spec(model, **opts)
|
|
257
|
+
# Full model-side parity: an explicit model_args dict merges over the
|
|
258
|
+
# convenience opts, so any lm-eval model arg (load_in_4bit, gptq, parallelize,
|
|
259
|
+
# max_memory, ...) is reachable — important for quantized/sharded models.
|
|
260
|
+
if model_args:
|
|
261
|
+
resolved_model_args.update(model_args)
|
|
262
|
+
|
|
263
|
+
try:
|
|
264
|
+
import lm_eval
|
|
265
|
+
except ImportError:
|
|
266
|
+
raise ExtraNotInstalled("lmeval", "pip install auditkit[lmeval]")
|
|
267
|
+
|
|
268
|
+
if backend in _CHAT_ONLY:
|
|
269
|
+
_assert_chat_compatible(tasks, model)
|
|
270
|
+
|
|
271
|
+
kwargs: dict[str, Any] = {
|
|
272
|
+
"model": backend,
|
|
273
|
+
"model_args": resolved_model_args,
|
|
274
|
+
"tasks": tasks,
|
|
275
|
+
"num_fewshot": cfg.num_fewshot,
|
|
276
|
+
"limit": cfg.limit,
|
|
277
|
+
"log_samples": True,
|
|
278
|
+
}
|
|
279
|
+
if cfg.batch_size is not None:
|
|
280
|
+
kwargs["batch_size"] = cfg.batch_size
|
|
281
|
+
if cfg.seed is not None:
|
|
282
|
+
kwargs["random_seed"] = cfg.seed
|
|
283
|
+
# Explicit run knobs (only forwarded when set, so defaults are unchanged).
|
|
284
|
+
for knob, value in (
|
|
285
|
+
("apply_chat_template", apply_chat_template),
|
|
286
|
+
("system_instruction", system_instruction),
|
|
287
|
+
("gen_kwargs", gen_kwargs),
|
|
288
|
+
("fewshot_as_multiturn", fewshot_as_multiturn),
|
|
289
|
+
):
|
|
290
|
+
if value is not None:
|
|
291
|
+
kwargs[knob] = value
|
|
292
|
+
# Full parity: any other simple_evaluate arg (max_batch_size, write_out,
|
|
293
|
+
# predict_only, task_manager, ...). Explicit args above win over this.
|
|
294
|
+
if lmeval_kwargs:
|
|
295
|
+
for k, v in lmeval_kwargs.items():
|
|
296
|
+
kwargs.setdefault(k, v)
|
|
297
|
+
|
|
298
|
+
fingerprint = _fingerprint(kwargs, model)
|
|
299
|
+
cache = DiskCache()
|
|
300
|
+
cached = cache.get(fingerprint)
|
|
301
|
+
if cached is not None:
|
|
302
|
+
return cached
|
|
303
|
+
|
|
304
|
+
# lm-eval constructs vLLM itself, not VLLMModel, so it never inherits
|
|
305
|
+
# vllm_gen's Colab-safe defaults. Without V1 multiprocessing off, engine
|
|
306
|
+
# core init fails in notebooks (empty Failed core proc(s)). Without the
|
|
307
|
+
# stdout fileno swap, vLLM's LLM() construction crashes in Colab
|
|
308
|
+
# (UnsupportedOperation: fileno on ipykernel's iostream).
|
|
309
|
+
stdout_cm: Any = contextlib.nullcontext()
|
|
310
|
+
if backend == "vllm":
|
|
311
|
+
from .model.vllm_gen import _apply_environment_defaults, _stdout_fix
|
|
312
|
+
_apply_environment_defaults()
|
|
313
|
+
stdout_cm = _stdout_fix()
|
|
314
|
+
|
|
315
|
+
hf_token = opts.get("hf_token") or opts.get("token")
|
|
316
|
+
# Groq/OpenRouter both authenticate the openai-chat-completions backend via
|
|
317
|
+
# OPENAI_API_KEY; source the key from api_key= or the provider's own env
|
|
318
|
+
# var and mirror it for the run only.
|
|
319
|
+
provider_key = None
|
|
320
|
+
if isinstance(model, str):
|
|
321
|
+
provider_prefix = model.split(":", 1)[0]
|
|
322
|
+
if provider_prefix == "groq":
|
|
323
|
+
provider_key = opts.get("api_key") or os.environ.get("GROQ_API_KEY")
|
|
324
|
+
if not provider_key:
|
|
325
|
+
raise AuditKitError(
|
|
326
|
+
"'groq:' benchmark run needs a Groq API key -- set GROQ_API_KEY or "
|
|
327
|
+
"pass api_key=. (Groq maps to lm-eval's openai-chat-completions "
|
|
328
|
+
"backend, which authenticates via OPENAI_API_KEY; AuditKIT mirrors "
|
|
329
|
+
"your Groq key into it for the duration of the run only.)"
|
|
330
|
+
)
|
|
331
|
+
elif provider_prefix == "openrouter":
|
|
332
|
+
provider_key = opts.get("api_key") or os.environ.get("OPENROUTER_API_KEY")
|
|
333
|
+
if not provider_key:
|
|
334
|
+
raise AuditKitError(
|
|
335
|
+
"'openrouter:' benchmark run needs an OpenRouter API key -- set "
|
|
336
|
+
"OPENROUTER_API_KEY or pass api_key=. (OpenRouter maps to lm-eval's "
|
|
337
|
+
"openai-chat-completions backend, which authenticates via "
|
|
338
|
+
"OPENAI_API_KEY; AuditKIT mirrors your OpenRouter key into it for "
|
|
339
|
+
"the duration of the run only.)"
|
|
340
|
+
)
|
|
341
|
+
try:
|
|
342
|
+
with _hf_token_env(hf_token), _temp_env("OPENAI_API_KEY", provider_key), stdout_cm:
|
|
343
|
+
raw = lm_eval.simple_evaluate(**kwargs)
|
|
344
|
+
except CapabilityError:
|
|
345
|
+
raise
|
|
346
|
+
except Exception as exc: # noqa: BLE001 — rewrap only the chat/MC case
|
|
347
|
+
if backend in _CHAT_ONLY and _looks_like_loglikelihood_error(exc):
|
|
348
|
+
raise CapabilityError(
|
|
349
|
+
f"task(s) {tasks} score by loglikelihood, but {model!r} is "
|
|
350
|
+
f"chat-only (lm-eval backend {backend!r}, no logprobs). Use a "
|
|
351
|
+
f"logprob-capable model (hf:/vllm:/api: with base_url)."
|
|
352
|
+
) from exc
|
|
353
|
+
raise
|
|
354
|
+
|
|
355
|
+
result = _to_runresult(raw, tasks, model, cfg, run_name, fingerprint)
|
|
356
|
+
cache.set(fingerprint, result)
|
|
357
|
+
return result
|
|
358
|
+
|
|
359
|
+
|
|
360
|
+
def _metric_bases(task_metrics: dict[str, Any]) -> list[str]:
|
|
361
|
+
"""The real metric base names from an lm-eval results dict, de-duplicated.
|
|
362
|
+
|
|
363
|
+
lm-eval keys real metrics as ``"{metric},{filter}"`` (``"acc,none"``,
|
|
364
|
+
``"exact_match,strict-match"``). Bookkeeping keys like ``"alias"`` and
|
|
365
|
+
``"sample_len"`` carry no comma, so requiring one drops them — otherwise
|
|
366
|
+
``sample_len`` (the sample count) leaks in as a bogus metric.
|
|
367
|
+
"""
|
|
368
|
+
bases: list[str] = []
|
|
369
|
+
for k in task_metrics:
|
|
370
|
+
if "," not in k:
|
|
371
|
+
continue
|
|
372
|
+
base = k.split(",")[0]
|
|
373
|
+
if base.endswith("_stderr") or base in bases:
|
|
374
|
+
continue
|
|
375
|
+
bases.append(base)
|
|
376
|
+
return bases
|
|
377
|
+
|
|
378
|
+
|
|
379
|
+
def _primary_metric(task_metrics: dict[str, Any]) -> str:
|
|
380
|
+
bases = _metric_bases(task_metrics)
|
|
381
|
+
for cand in _KNOWN_SAMPLE_METRICS:
|
|
382
|
+
if cand in bases:
|
|
383
|
+
return cand
|
|
384
|
+
return bases[0] if bases else "acc"
|
|
385
|
+
|
|
386
|
+
|
|
387
|
+
def _extract_prompt(sample: dict) -> str:
|
|
388
|
+
args = sample.get("arguments")
|
|
389
|
+
if args:
|
|
390
|
+
first = args[0]
|
|
391
|
+
if isinstance(first, (list, tuple)) and first:
|
|
392
|
+
return str(first[0])
|
|
393
|
+
if isinstance(first, dict):
|
|
394
|
+
for v in first.values():
|
|
395
|
+
return str(v)
|
|
396
|
+
return str(first)
|
|
397
|
+
doc = sample.get("doc")
|
|
398
|
+
if isinstance(doc, dict):
|
|
399
|
+
for key in ("question", "query", "input", "text", "ctx", "goal"):
|
|
400
|
+
if key in doc:
|
|
401
|
+
return str(doc[key])
|
|
402
|
+
return str(doc) if doc is not None else ""
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
def _as_number(x: Any) -> Optional[float]:
|
|
406
|
+
if isinstance(x, bool):
|
|
407
|
+
return None
|
|
408
|
+
if isinstance(x, (int, float)):
|
|
409
|
+
return float(x)
|
|
410
|
+
if isinstance(x, (list, tuple)) and x and isinstance(x[0], (int, float)) and not isinstance(x[0], bool):
|
|
411
|
+
return float(x[0]) # lm-eval logs MC choices as [loglikelihood, is_greedy]
|
|
412
|
+
return None
|
|
413
|
+
|
|
414
|
+
|
|
415
|
+
def _extract_output(sample: dict) -> str:
|
|
416
|
+
"""The model's answer, made readable.
|
|
417
|
+
|
|
418
|
+
For multiple-choice tasks lm-eval's ``filtered_resps`` are per-choice
|
|
419
|
+
loglikelihoods, so the raw first element is a meaningless number; we report
|
|
420
|
+
the argmax **choice index** instead (which lines up with the gold ``target``
|
|
421
|
+
index). For generative tasks we report the response text as-is.
|
|
422
|
+
"""
|
|
423
|
+
for key in ("filtered_resps", "resps"):
|
|
424
|
+
r = sample.get(key)
|
|
425
|
+
if not r:
|
|
426
|
+
continue
|
|
427
|
+
nums = [_as_number(item) for item in r]
|
|
428
|
+
if len(nums) > 1 and all(n is not None for n in nums):
|
|
429
|
+
return str(max(range(len(nums)), key=lambda i: nums[i])) # MC → picked index
|
|
430
|
+
first = r[0] if isinstance(r, (list, tuple)) else r
|
|
431
|
+
if isinstance(first, (list, tuple)) and first:
|
|
432
|
+
return str(first[0])
|
|
433
|
+
return str(first)
|
|
434
|
+
return ""
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
def _sample_to_prediction(
|
|
438
|
+
sample: dict, task: str, run_id: str, idx: int, primary_base: str
|
|
439
|
+
) -> Prediction:
|
|
440
|
+
target = sample.get("target")
|
|
441
|
+
raw_score = sample.get(primary_base)
|
|
442
|
+
score = float(raw_score) if isinstance(raw_score, (int, float, bool)) else None
|
|
443
|
+
return Prediction(
|
|
444
|
+
run_id=run_id,
|
|
445
|
+
task=task,
|
|
446
|
+
sample_id=str(sample.get("doc_id", idx)),
|
|
447
|
+
prompt=_extract_prompt(sample),
|
|
448
|
+
raw_output=_extract_output(sample),
|
|
449
|
+
parsed_answer=_extract_output(sample),
|
|
450
|
+
expected=None if target is None else str(target),
|
|
451
|
+
correct=(score == 1.0) if score is not None else None,
|
|
452
|
+
score=score,
|
|
453
|
+
metadata={"engine": "lmeval", "primary_metric": primary_base},
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
|
|
457
|
+
def _fingerprint(kwargs: dict, model_spec: str) -> str:
|
|
458
|
+
"""Identity of a benchmark run, computed from the exact kwargs about to be
|
|
459
|
+
passed to ``lm_eval.simple_evaluate`` -- so it can be checked *before*
|
|
460
|
+
running (cache lookup) and reused as-is for the cached :class:`RunResult`,
|
|
461
|
+
the same way :meth:`Runner.run` fingerprints a native ``RunSpec`` before
|
|
462
|
+
executing it.
|
|
463
|
+
"""
|
|
464
|
+
try:
|
|
465
|
+
import lm_eval
|
|
466
|
+
|
|
467
|
+
version = getattr(lm_eval, "__version__", "?")
|
|
468
|
+
except Exception:
|
|
469
|
+
version = "?"
|
|
470
|
+
key = {
|
|
471
|
+
"engine": "lmeval",
|
|
472
|
+
"model_spec": model_spec,
|
|
473
|
+
"lm_eval_version": version,
|
|
474
|
+
**{k: (sorted(v) if k == "tasks" and isinstance(v, list) else v) for k, v in kwargs.items()},
|
|
475
|
+
}
|
|
476
|
+
blob = json.dumps(key, sort_keys=True, default=str).encode("utf-8")
|
|
477
|
+
return hashlib.sha256(blob).hexdigest()[:16]
|
|
478
|
+
|
|
479
|
+
|
|
480
|
+
def _to_runresult(
|
|
481
|
+
raw: dict, tasks: list[str], model_spec: str, cfg: Any, run_name: str, fingerprint: str
|
|
482
|
+
) -> RunResult:
|
|
483
|
+
"""Map lm-eval's ``simple_evaluate`` output to a :class:`RunResult`."""
|
|
484
|
+
results = raw.get("results", {}) or {}
|
|
485
|
+
samples_by_task = raw.get("samples", {}) or {}
|
|
486
|
+
|
|
487
|
+
run_id = run_name or f"benchmark-{fingerprint}"
|
|
488
|
+
|
|
489
|
+
# Headline: the authoritative lm-eval aggregate numbers. Real metrics are
|
|
490
|
+
# keyed "{metric},{filter}"; drop the noise "none" filter for a clean name
|
|
491
|
+
# but keep any other filter (e.g. gsm8k strict-match vs flexible-extract).
|
|
492
|
+
headline: dict[str, float] = {}
|
|
493
|
+
for task, metrics in results.items():
|
|
494
|
+
for key, val in metrics.items():
|
|
495
|
+
if "," not in key or not isinstance(val, (int, float, bool)):
|
|
496
|
+
continue
|
|
497
|
+
metric, filt = key.split(",", 1)
|
|
498
|
+
if metric.endswith("_stderr"):
|
|
499
|
+
continue
|
|
500
|
+
name = f"{task}:{metric}" if filt == "none" else f"{task}:{metric},{filt}"
|
|
501
|
+
headline[name] = float(val)
|
|
502
|
+
|
|
503
|
+
# Per-sample: one Prediction per logged doc; per-metric Stats from real values.
|
|
504
|
+
stats: dict[str, Stat] = {}
|
|
505
|
+
predictions: list[Prediction] = []
|
|
506
|
+
for task, samps in samples_by_task.items():
|
|
507
|
+
task_metrics = results.get(task, {})
|
|
508
|
+
primary_base = _primary_metric(task_metrics)
|
|
509
|
+
metric_bases = _metric_bases(task_metrics)
|
|
510
|
+
for idx, s in enumerate(samps):
|
|
511
|
+
predictions.append(_sample_to_prediction(s, task, run_id, idx, primary_base))
|
|
512
|
+
for base in metric_bases:
|
|
513
|
+
v = s.get(base)
|
|
514
|
+
if isinstance(v, (int, float, bool)):
|
|
515
|
+
stats.setdefault(f"{task}:{base}", Stat(f"{task}:{base}")).add(float(v))
|
|
516
|
+
|
|
517
|
+
# Any headline metric with no per-sample values still gets a single-value Stat.
|
|
518
|
+
for name, val in headline.items():
|
|
519
|
+
if name not in stats:
|
|
520
|
+
stats[name] = Stat(name).add(val)
|
|
521
|
+
|
|
522
|
+
return RunResult(
|
|
523
|
+
run_id=run_id,
|
|
524
|
+
fingerprint=fingerprint,
|
|
525
|
+
stats=stats,
|
|
526
|
+
predictions=predictions,
|
|
527
|
+
headline=headline,
|
|
528
|
+
config=cfg,
|
|
529
|
+
model_spec=model_spec,
|
|
530
|
+
)
|
|
531
|
+
|
|
532
|
+
|
|
533
|
+
class BenchmarkEvaluator:
|
|
534
|
+
"""The lm-eval technique engine as a reusable, named object.
|
|
535
|
+
|
|
536
|
+
Thin wrapper over :func:`run_benchmark` so a benchmark run can be composed
|
|
537
|
+
and passed around. Not a :class:`~auditkit.model.Model` — it owns task
|
|
538
|
+
loading and scoring, and produces a full :class:`RunResult` via
|
|
539
|
+
:meth:`run`.
|
|
540
|
+
"""
|
|
541
|
+
|
|
542
|
+
technique = "benchmark"
|
|
543
|
+
|
|
544
|
+
def __init__(self, tasks: Union[str, list[str]], **opts: Any) -> None:
|
|
545
|
+
self.tasks = [tasks] if isinstance(tasks, str) else list(tasks)
|
|
546
|
+
self.opts = opts
|
|
547
|
+
|
|
548
|
+
def run(self, model: Any, *, config: Any = None, run_name: str = "", **opts: Any) -> RunResult:
|
|
549
|
+
merged = {**self.opts, **opts}
|
|
550
|
+
return run_benchmark(self.tasks, model, config=config, run_name=run_name, **merged)
|
auditkit/loaders.py
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
"""Data loaders that produce lists of :class:`~auditkit.sample.Sample`.
|
|
2
|
+
|
|
3
|
+
T0 provides ``load_csv`` (stdlib only). Optional loaders (``load_hf``,
|
|
4
|
+
``load_croissant``) are deferred to their own cycles and raise
|
|
5
|
+
:class:`ExtraNotInstalled` when their extra is absent.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import csv
|
|
11
|
+
from typing import Any, Optional
|
|
12
|
+
|
|
13
|
+
from .errors import ExtraNotInstalled
|
|
14
|
+
from .sample import Sample
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def load_csv(
|
|
18
|
+
path: str,
|
|
19
|
+
*,
|
|
20
|
+
input_col: str = "input",
|
|
21
|
+
target_col: Optional[str] = "target",
|
|
22
|
+
output_col: Optional[str] = None,
|
|
23
|
+
**kwargs: Any,
|
|
24
|
+
) -> list[Sample]:
|
|
25
|
+
"""Load samples from a CSV file.
|
|
26
|
+
|
|
27
|
+
Parameters
|
|
28
|
+
----------
|
|
29
|
+
path : str
|
|
30
|
+
Path to the CSV file.
|
|
31
|
+
input_col : str
|
|
32
|
+
Column name for the sample input (default ``"input"``).
|
|
33
|
+
target_col : str or None
|
|
34
|
+
Column name for the sample target (default ``"target"``);
|
|
35
|
+
``None`` means no target column is expected.
|
|
36
|
+
output_col : str or None
|
|
37
|
+
Column name for a pre-generated answer (default ``None``); when set, it
|
|
38
|
+
fills ``Sample.actual_output`` so the rows can be scored directly with
|
|
39
|
+
``model="precomputed"`` (no generation step).
|
|
40
|
+
**kwargs
|
|
41
|
+
Extra keyword arguments forwarded to ``csv.DictReader``.
|
|
42
|
+
"""
|
|
43
|
+
with open(path, newline="", encoding="utf-8") as fh:
|
|
44
|
+
reader = csv.DictReader(fh, **kwargs)
|
|
45
|
+
samples: list[Sample] = []
|
|
46
|
+
for row in reader:
|
|
47
|
+
target = row.get(target_col) if target_col else None # type: ignore[arg-type]
|
|
48
|
+
output = row.get(output_col) if output_col else None
|
|
49
|
+
samples.append(
|
|
50
|
+
Sample(
|
|
51
|
+
input=row[input_col],
|
|
52
|
+
target=target,
|
|
53
|
+
actual_output=output,
|
|
54
|
+
)
|
|
55
|
+
)
|
|
56
|
+
return samples
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def load_hf(
|
|
60
|
+
path: str,
|
|
61
|
+
*,
|
|
62
|
+
split: str = "test",
|
|
63
|
+
input_col: str = "question",
|
|
64
|
+
target_col: str = "answer",
|
|
65
|
+
**kwargs: Any,
|
|
66
|
+
) -> list[Sample]:
|
|
67
|
+
try:
|
|
68
|
+
import datasets
|
|
69
|
+
except ImportError:
|
|
70
|
+
raise ExtraNotInstalled("interop", "HuggingFace datasets (install auditkit[interop])") from None
|
|
71
|
+
records = datasets.load_dataset(path, split=split, **kwargs)
|
|
72
|
+
return [Sample(input=row[input_col], target=row.get(target_col)) for row in records]
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def load_croissant(
|
|
76
|
+
path: str,
|
|
77
|
+
record_set: str,
|
|
78
|
+
*,
|
|
79
|
+
input_col: str = "input",
|
|
80
|
+
target_col: str = "target",
|
|
81
|
+
**kwargs: Any,
|
|
82
|
+
) -> list[Sample]:
|
|
83
|
+
"""Load samples from a Croissant (JSON-LD) dataset descriptor.
|
|
84
|
+
|
|
85
|
+
Parameters
|
|
86
|
+
----------
|
|
87
|
+
path : str
|
|
88
|
+
Path or URL to the Croissant JSON-LD file (e.g. a HuggingFace
|
|
89
|
+
dataset's ``.../croissant`` endpoint).
|
|
90
|
+
record_set : str
|
|
91
|
+
The Croissant record set to read (a Croissant file can describe
|
|
92
|
+
several; there's no universally correct default). Inspect
|
|
93
|
+
``mlcroissant.Dataset(jsonld=path).metadata.record_sets`` to see
|
|
94
|
+
what's available for a given dataset.
|
|
95
|
+
input_col : str
|
|
96
|
+
Field name for the sample input, *without* the record-set prefix
|
|
97
|
+
Croissant adds (e.g. ``"question"``, not ``"question-answer/question"``
|
|
98
|
+
-- the prefix is added automatically).
|
|
99
|
+
target_col : str
|
|
100
|
+
Field name for the sample target, same convention as ``input_col``.
|
|
101
|
+
**kwargs
|
|
102
|
+
Extra keyword arguments forwarded to ``mlcroissant.Dataset(...)``.
|
|
103
|
+
"""
|
|
104
|
+
try:
|
|
105
|
+
import mlcroissant
|
|
106
|
+
except ImportError:
|
|
107
|
+
raise ExtraNotInstalled("interop", "mlcroissant (install auditkit[interop])") from None
|
|
108
|
+
|
|
109
|
+
def _decode(value: Any) -> Any:
|
|
110
|
+
return value.decode("utf-8") if isinstance(value, bytes) else value
|
|
111
|
+
|
|
112
|
+
dataset = mlcroissant.Dataset(jsonld=path, **kwargs)
|
|
113
|
+
input_key = f"{record_set}/{input_col}"
|
|
114
|
+
target_key = f"{record_set}/{target_col}"
|
|
115
|
+
samples = []
|
|
116
|
+
for row in dataset.records(record_set):
|
|
117
|
+
samples.append(Sample(
|
|
118
|
+
input=_decode(row[input_key]),
|
|
119
|
+
target=_decode(row.get(target_key)),
|
|
120
|
+
))
|
|
121
|
+
return samples
|
auditkit/logs.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""Logging setup for AuditKIT."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import logging
|
|
5
|
+
import sys
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def configure_logging(*, level: int | str = logging.INFO, fmt: str | None = None) -> None:
|
|
9
|
+
"""Configure auditkit logging with a structured format."""
|
|
10
|
+
logger = logging.getLogger("auditkit")
|
|
11
|
+
logger.setLevel(level)
|
|
12
|
+
if not logger.handlers:
|
|
13
|
+
handler = logging.StreamHandler(sys.stderr)
|
|
14
|
+
handler.setFormatter(logging.Formatter(
|
|
15
|
+
fmt or "%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
|
16
|
+
datefmt="%H:%M:%S",
|
|
17
|
+
))
|
|
18
|
+
logger.addHandler(handler)
|