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.
Files changed (81) hide show
  1. auditkit/README.md +99 -0
  2. auditkit/__init__.py +177 -0
  3. auditkit/__main__.py +3 -0
  4. auditkit/_bootstrap.py +77 -0
  5. auditkit/_identity_guard.py +99 -0
  6. auditkit/adapter.py +264 -0
  7. auditkit/annotator.py +339 -0
  8. auditkit/api.py +502 -0
  9. auditkit/assets/auditkit_logo.png +0 -0
  10. auditkit/cache.py +47 -0
  11. auditkit/cli.py +417 -0
  12. auditkit/comparison.py +563 -0
  13. auditkit/diff.py +265 -0
  14. auditkit/errors.py +54 -0
  15. auditkit/evaluator.py +20 -0
  16. auditkit/experiment.py +145 -0
  17. auditkit/hf_publish.py +262 -0
  18. auditkit/lmeval_engine.py +550 -0
  19. auditkit/loaders.py +121 -0
  20. auditkit/logs.py +18 -0
  21. auditkit/metric.py +199 -0
  22. auditkit/metrics/README.md +15 -0
  23. auditkit/metrics/__init__.py +0 -0
  24. auditkit/metrics/code.py +222 -0
  25. auditkit/metrics/embedding.py +131 -0
  26. auditkit/metrics/encoder_judge.py +423 -0
  27. auditkit/metrics/generation.py +331 -0
  28. auditkit/metrics/guard.py +412 -0
  29. auditkit/metrics/hallucination.py +45 -0
  30. auditkit/metrics/judge.py +547 -0
  31. auditkit/metrics/pairwise.py +153 -0
  32. auditkit/metrics/perf.py +53 -0
  33. auditkit/metrics/rag.py +149 -0
  34. auditkit/metrics/security.py +64 -0
  35. auditkit/metrics/toxicity.py +238 -0
  36. auditkit/model/README.md +16 -0
  37. auditkit/model/__init__.py +485 -0
  38. auditkit/model/anthropic.py +94 -0
  39. auditkit/model/api_gen.py +133 -0
  40. auditkit/model/groq_gen.py +121 -0
  41. auditkit/model/hf_gen.py +385 -0
  42. auditkit/model/lexsi.py +155 -0
  43. auditkit/model/litellm_gen.py +65 -0
  44. auditkit/model/openai.py +90 -0
  45. auditkit/model/openrouter_gen.py +152 -0
  46. auditkit/model/vllm_gen.py +316 -0
  47. auditkit/model_compare.py +655 -0
  48. auditkit/redteam/README.md +9 -0
  49. auditkit/redteam/__init__.py +26 -0
  50. auditkit/redteam/detector.py +37 -0
  51. auditkit/redteam/detectors/README.md +5 -0
  52. auditkit/redteam/detectors/builtin.py +126 -0
  53. auditkit/redteam/probe.py +39 -0
  54. auditkit/redteam/probes/README.md +5 -0
  55. auditkit/redteam/probes/builtin.py +85 -0
  56. auditkit/redteam/runner.py +206 -0
  57. auditkit/registry.py +65 -0
  58. auditkit/report.py +278 -0
  59. auditkit/report_format.py +52 -0
  60. auditkit/router.py +54 -0
  61. auditkit/runner.py +575 -0
  62. auditkit/runspec.py +159 -0
  63. auditkit/sample.py +40 -0
  64. auditkit/scenario.py +88 -0
  65. auditkit/scenarios/README.md +10 -0
  66. auditkit/scenarios/__init__.py +4 -0
  67. auditkit/scenarios/arc.py +33 -0
  68. auditkit/scenarios/gsm8k.py +32 -0
  69. auditkit/scenarios/hellaswag.py +33 -0
  70. auditkit/scenarios/humaneval.py +32 -0
  71. auditkit/scenarios/mmlu.py +34 -0
  72. auditkit/scenarios/truthfulqa.py +33 -0
  73. auditkit/score.py +165 -0
  74. auditkit/scorers.py +117 -0
  75. auditkit/scoring.py +79 -0
  76. auditkit/types.py +69 -0
  77. auditkit-1.0.0.dist-info/METADATA +396 -0
  78. auditkit-1.0.0.dist-info/RECORD +81 -0
  79. auditkit-1.0.0.dist-info/WHEEL +4 -0
  80. auditkit-1.0.0.dist-info/entry_points.txt +2 -0
  81. 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)