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
auditkit/metric.py ADDED
@@ -0,0 +1,199 @@
1
+ """Metrics: the internal, richer form of a scorer.
2
+
3
+ A :class:`Metric` turns one sample plus one model output into one or more
4
+ :class:`Score`. It knows which sample fields it needs (``required_fields``), so
5
+ the runner can skip a metric that doesn't apply instead of crashing — an
6
+ exact-match metric is silently skipped for an open-ended sample with no gold.
7
+ The two here are deterministic string matchers; judge/code/security metrics
8
+ arrive as later techniques but wear this same interface.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import re
14
+ from abc import ABC, abstractmethod
15
+ from typing import Any, Optional, Union
16
+
17
+ from .registry import METRICS
18
+ from .sample import Sample
19
+ from .score import Score
20
+ from .types import Direction, ScoreKind
21
+ from ._identity_guard import warn_if_identity_incomplete
22
+
23
+
24
+ class Metric(ABC):
25
+ """Score a model output against a sample."""
26
+
27
+ name: str = "metric"
28
+ kind: ScoreKind = ScoreKind.BENCHMARK
29
+ is_deterministic: bool = True
30
+ required_fields: frozenset[str] = frozenset()
31
+
32
+ def __init_subclass__(cls, **kwargs) -> None:
33
+ super().__init_subclass__(**kwargs)
34
+ warn_if_identity_incomplete(cls, Metric, "score")
35
+ if not hasattr(cls, "direction"):
36
+ raise TypeError(
37
+ f"{cls.__name__} must declare 'direction' (Direction.MAXIMIZE or "
38
+ f"Direction.MINIMIZE) -- required so RunComparison/compare_models() "
39
+ f"grade it in the right sense (a metric where lower is better, e.g. "
40
+ f"latency or error rate, graded as MAXIMIZE would report a regression "
41
+ f"as an improvement). Set it as a class attribute:\n\n"
42
+ f" class {cls.__name__}(Metric):\n"
43
+ f" direction = Direction.MAXIMIZE # or Direction.MINIMIZE\n\n"
44
+ f"See docs/best_practices/scorers.md#direction for a full example."
45
+ )
46
+
47
+ def applicable(self, sample: Sample) -> bool:
48
+ """False when any required field is absent on the sample (→ skipped)."""
49
+ for field_name in self.required_fields:
50
+ if getattr(sample, field_name, None) is None:
51
+ return False
52
+ return True
53
+
54
+ @abstractmethod
55
+ def score(self, sample: Sample, output: str, context: Any = None) -> Union[Score, list[Score]]:
56
+ """Produce a Score (or list) for ``output`` against ``sample``."""
57
+ ...
58
+
59
+ async def ascore(self, sample: Sample, output: str, context: Any = None):
60
+ """Async entry point; deterministic metrics just defer to :meth:`score`."""
61
+ return self.score(sample, output, context)
62
+
63
+ def identity(self) -> dict:
64
+ """The config that defines this metric's scoring, for the run fingerprint.
65
+
66
+ Deterministic metrics are fully described by their name. Judges override
67
+ this to include the judge model + prompt + choices, so changing the judge
68
+ (a different "ruler") changes the run fingerprint and never silently
69
+ reuses a cached result scored by a different judge. ``__init_subclass__``
70
+ above warns at class-definition time when a subclass takes constructor
71
+ arguments but skips this override (unless it instead makes ``self.name``
72
+ itself parameter-derived, which already protects the fingerprint — see
73
+ ``_identity_guard.py``)."""
74
+ return {"name": self.name}
75
+
76
+
77
+ @METRICS.register("exact_match")
78
+ class ExactMatch(Metric):
79
+ """1.0 iff the output equals the target after trimming outer whitespace."""
80
+
81
+ name = "exact_match"
82
+ direction = Direction.MAXIMIZE
83
+ required_fields = frozenset({"target"})
84
+
85
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
86
+ hit = output.strip() == (sample.target or "").strip()
87
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
88
+
89
+
90
+ _ARTICLE_RE = re.compile(r"\b(a|an|the)\b", re.IGNORECASE)
91
+ _PUNCT_RE = re.compile(r"[^\w\s]")
92
+
93
+
94
+ def _normalize(text: str) -> str:
95
+ """Lowercase, drop punctuation and articles, collapse whitespace."""
96
+ text = text.lower()
97
+ text = _PUNCT_RE.sub(" ", text)
98
+ text = _ARTICLE_RE.sub(" ", text)
99
+ return " ".join(text.split())
100
+
101
+
102
+ @METRICS.register("quasi_exact_match")
103
+ class QuasiExactMatch(Metric):
104
+ """Exact match after normalization (case, punctuation, articles, spacing)."""
105
+
106
+ name = "quasi_exact_match"
107
+ direction = Direction.MAXIMIZE
108
+ required_fields = frozenset({"target"})
109
+
110
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
111
+ hit = _normalize(output) == _normalize(sample.target or "")
112
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
113
+
114
+
115
+ def _resolve_choice_index(value: Any, choices: list) -> Optional[int]:
116
+ """Resolve a target or model-output value to a 0-based index into ``choices``.
117
+
118
+ ``value`` legitimately shows up in three encodings across scenarios and
119
+ custom loaders: a bare letter ("B", only meaningful for <=26 choices), a
120
+ 0-based index into ``choices`` (``1`` or ``"1"``), or the literal text of
121
+ the correct choice. All three are accepted here so ``Acc``/``AccNorm``
122
+ score correctly regardless of which one a given dataset uses, instead of
123
+ silently mis-scoring (or crashing on a non-str target, e.g. an int).
124
+
125
+ The canonical return value is a 0-based **index**, not a letter — letters
126
+ run out at 26 choices (``chr(65+26)`` is ``'['``, not a letter) and
127
+ collide once you reach lowercase range (index 32 -> ``'a'``, which
128
+ case-normalizes to the same key as index 0's ``'A'``). An index has no
129
+ such ceiling, so it's what both the adapter's prompt labels (see
130
+ :class:`~auditkit.adapter.MCQAdapter`) and this resolver key off of.
131
+ """
132
+ if value is None:
133
+ return None
134
+ choices = choices or []
135
+ if isinstance(value, str):
136
+ s = value.strip()
137
+ # Compare against *stripped* choices, not raw ones -- MCQ choices
138
+ # commonly carry a leading space (the correct convention for
139
+ # GPT-2/BPE-style loglikelihood scoring, so " Paris" tokenizes as a
140
+ # proper word-initial continuation of the prompt), and comparing an
141
+ # already-stripped candidate against unstripped choices can never
142
+ # match even on a perfect pick.
143
+ for i, c in enumerate(choices):
144
+ if isinstance(c, str) and c.strip() == s:
145
+ return i
146
+ if s.lstrip("-").isdigit():
147
+ idx = int(s)
148
+ if 0 <= idx < len(choices):
149
+ return idx
150
+ if len(s) == 1 and s.isalpha():
151
+ idx = ord(s.upper()) - ord("A")
152
+ if 0 <= idx < len(choices):
153
+ return idx
154
+ elif isinstance(value, int) and not isinstance(value, bool):
155
+ if 0 <= value < len(choices):
156
+ return value
157
+ return None
158
+
159
+
160
+ def _target_index(sample: Sample) -> int:
161
+ idx = _resolve_choice_index(sample.target, sample.choices or [])
162
+ if idx is None:
163
+ raise ValueError(
164
+ f"cannot resolve sample.target={sample.target!r} against "
165
+ f"sample.choices={sample.choices!r}; expected a 0-based index, "
166
+ f"a choice letter (e.g. 'B', only for <=26 choices), or the "
167
+ f"exact text of the correct choice"
168
+ )
169
+ return idx
170
+
171
+
172
+ @METRICS.register("acc")
173
+ class Acc(Metric):
174
+ """1.0 iff the output resolves to the same choice index as the target."""
175
+
176
+ name = "acc"
177
+ direction = Direction.MAXIMIZE
178
+ required_fields = frozenset({"target", "choices"})
179
+
180
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
181
+ target_idx = _target_index(sample)
182
+ pred_idx = _resolve_choice_index(output, sample.choices or [])
183
+ hit = pred_idx is not None and pred_idx == target_idx
184
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
185
+
186
+
187
+ @METRICS.register("acc_norm")
188
+ class AccNorm(Metric):
189
+ """1.0 iff the predicted choice text maps to the correct target choice."""
190
+
191
+ name = "acc_norm"
192
+ direction = Direction.MAXIMIZE
193
+ required_fields = frozenset({"target", "choices"})
194
+
195
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
196
+ target_idx = _target_index(sample)
197
+ pred_idx = _resolve_choice_index(output, sample.choices or [])
198
+ hit = pred_idx is not None and pred_idx == target_idx
199
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
@@ -0,0 +1,15 @@
1
+ # Metrics
2
+
3
+ Built-in metric families for evaluating model outputs:
4
+
5
+ - **code** — Contains, Equals, F1Score, IsJson, Levenshtein, Regex, StartsWith, EndsWith, WordCount
6
+ - **embedding** — CosineSimilarity, TokenOverlap, BM25Similarity
7
+ - **generation** — Bleu, RogueL, ChrF, BertScore, Perplexity, WordErrorRate
8
+ - **hallucination** — FactualConsistency
9
+ - **judge** — JudgeMetric, LLMJudge, GEval, RubricItem, Factuality, ClosedQA, Relevance, BiasJudge
10
+ - **pairwise** — WinRate, EloScore, PreferenceAccuracy
11
+ - **perf** — LatencyStats, Throughput
12
+ - **rag** — LexicalGroundedness, ContextCoverage, ContextOverlap, AnswerOverlap
13
+ - **security** — DefconGrade, KeywordDetector, ThreatCategory
14
+ - **toxicity** — ToxicityScore, RepresentationSkew, HateSpeechScore
15
+ - **guard** — GuardJudge (safety scoring via a guard model, e.g. Llama Guard)
File without changes
@@ -0,0 +1,222 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import re
5
+ from typing import Any
6
+
7
+ from auditkit.metric import Metric
8
+ from auditkit.registry import METRICS
9
+ from auditkit.sample import Sample
10
+ from auditkit.score import Score
11
+ from auditkit.types import Direction, ScoreKind
12
+
13
+
14
+ @METRICS.register("equals")
15
+ class Equals(Metric):
16
+ kind = ScoreKind.CODE
17
+ direction = Direction.MAXIMIZE
18
+
19
+ def __init__(self, ignore_case: bool = False) -> None:
20
+ self._ignore_case = ignore_case
21
+ self.name = "equals_ci" if ignore_case else "equals"
22
+ self.required_fields = frozenset({"target"})
23
+
24
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
25
+ a = output.strip()
26
+ b = (sample.target or "").strip()
27
+ if self._ignore_case:
28
+ hit = a.lower() == b.lower()
29
+ else:
30
+ hit = a == b
31
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
32
+
33
+
34
+ @METRICS.register("contains")
35
+ class Contains(Metric):
36
+ kind = ScoreKind.CODE
37
+ direction = Direction.MAXIMIZE
38
+
39
+ def __init__(self, substring: str, ignore_case: bool = False) -> None:
40
+ self._substring = substring
41
+ self._ignore_case = ignore_case
42
+ prefix = "contains_ci" if ignore_case else "contains"
43
+ self.name = f"{prefix}({substring})"
44
+ self.required_fields = frozenset()
45
+
46
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
47
+ if self._ignore_case:
48
+ hit = self._substring.lower() in output.lower()
49
+ else:
50
+ hit = self._substring in output
51
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
52
+
53
+
54
+ @METRICS.register("starts_with")
55
+ class StartsWith(Metric):
56
+ kind = ScoreKind.CODE
57
+ direction = Direction.MAXIMIZE
58
+
59
+ def __init__(self, prefix: str, ignore_case: bool = False) -> None:
60
+ self._prefix = prefix
61
+ self._ignore_case = ignore_case
62
+ prefix_label = "startswith_ci" if ignore_case else "startswith"
63
+ self.name = f"{prefix_label}({prefix})"
64
+ self.required_fields = frozenset()
65
+
66
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
67
+ if self._ignore_case:
68
+ hit = output.lower().startswith(self._prefix.lower())
69
+ else:
70
+ hit = output.startswith(self._prefix)
71
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
72
+
73
+
74
+ @METRICS.register("ends_with")
75
+ class EndsWith(Metric):
76
+ kind = ScoreKind.CODE
77
+ direction = Direction.MAXIMIZE
78
+
79
+ def __init__(self, suffix: str, ignore_case: bool = False) -> None:
80
+ self._suffix = suffix
81
+ self._ignore_case = ignore_case
82
+ self.name = f"endswith({suffix})"
83
+ self.required_fields = frozenset()
84
+
85
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
86
+ if self._ignore_case:
87
+ hit = output.lower().endswith(self._suffix.lower())
88
+ else:
89
+ hit = output.endswith(self._suffix)
90
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
91
+
92
+
93
+ @METRICS.register("regex")
94
+ class Regex(Metric):
95
+ kind = ScoreKind.CODE
96
+ direction = Direction.MAXIMIZE
97
+
98
+ def __init__(self, pattern: str) -> None:
99
+ self._compiled = re.compile(pattern)
100
+ self.name = f"regex({pattern})"
101
+ self.required_fields = frozenset()
102
+
103
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
104
+ hit = bool(self._compiled.search(output))
105
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
106
+
107
+
108
+ @METRICS.register("levenshtein")
109
+ class Levenshtein(Metric):
110
+ kind = ScoreKind.CODE
111
+ direction = Direction.MAXIMIZE
112
+
113
+ def __init__(self) -> None:
114
+ self.name = "levenshtein"
115
+ self.required_fields = frozenset()
116
+
117
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
118
+ ref = sample.target or ""
119
+ if not output and not ref:
120
+ return Score(name=self.name, value=0.0, kind=self.kind)
121
+ distance = self._edit_distance(output, ref)
122
+ max_len = max(len(output), len(ref))
123
+ score_val = 1.0 - distance / max_len
124
+ return Score(name=self.name, value=score_val, kind=self.kind)
125
+
126
+ @staticmethod
127
+ def _edit_distance(a: str, b: str) -> int:
128
+ m, n = len(a), len(b)
129
+ prev = list(range(n + 1))
130
+ for i in range(1, m + 1):
131
+ curr = [i] * (n + 1)
132
+ for j in range(1, n + 1):
133
+ cost = 0 if a[i - 1] == b[j - 1] else 1
134
+ curr[j] = min(
135
+ prev[j] + 1,
136
+ curr[j - 1] + 1,
137
+ prev[j - 1] + cost,
138
+ )
139
+ prev = curr
140
+ return prev[n]
141
+
142
+
143
+ @METRICS.register("word_count")
144
+ class WordCount(Metric):
145
+ kind = ScoreKind.CODE
146
+ direction = Direction.MAXIMIZE
147
+
148
+ def __init__(self, min_words: int = 0, max_words: int | None = None) -> None:
149
+ self._min = min_words
150
+ self._max = max_words
151
+ if min_words > 0 and max_words is not None:
152
+ self.name = f"wordcount({min_words}-{max_words})"
153
+ elif min_words > 0:
154
+ self.name = f"wordcount({min_words}-)"
155
+ elif max_words is not None:
156
+ self.name = f"wordcount(-{max_words})"
157
+ else:
158
+ self.name = f"wordcount({min_words}-)"
159
+ self.required_fields = frozenset()
160
+
161
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
162
+ count = len(output.split())
163
+ if count < self._min:
164
+ hit = False
165
+ elif self._max is not None and count > self._max:
166
+ hit = False
167
+ else:
168
+ hit = True
169
+ return Score(name=self.name, value=1.0 if hit else 0.0, kind=self.kind)
170
+
171
+
172
+ @METRICS.register("is_json")
173
+ class IsJson(Metric):
174
+ kind = ScoreKind.CODE
175
+ direction = Direction.MAXIMIZE
176
+
177
+ def __init__(self, require_keys: list[str] | None = None) -> None:
178
+ self._require_keys = require_keys
179
+ if require_keys:
180
+ keys_str = ",".join(require_keys)
181
+ self.name = f"is_json(keys={keys_str})"
182
+ else:
183
+ self.name = "is_json"
184
+ self.required_fields = frozenset()
185
+
186
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
187
+ try:
188
+ parsed = json.loads(output)
189
+ except (json.JSONDecodeError, ValueError):
190
+ return Score(name=self.name, value=0.0, kind=self.kind)
191
+ if self._require_keys is not None:
192
+ if not isinstance(parsed, dict):
193
+ return Score(name=self.name, value=0.0, kind=self.kind)
194
+ if not all(k in parsed for k in self._require_keys):
195
+ return Score(name=self.name, value=0.0, kind=self.kind)
196
+ return Score(name=self.name, value=1.0, kind=self.kind)
197
+
198
+
199
+ @METRICS.register("f1_score")
200
+ class F1Score(Metric):
201
+ kind = ScoreKind.CODE
202
+ direction = Direction.MAXIMIZE
203
+
204
+ def __init__(self) -> None:
205
+ self.name = "f1_score"
206
+ self.required_fields = frozenset({"target"})
207
+
208
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
209
+ output_tokens = output.split()
210
+ target_tokens = (sample.target or "").split()
211
+ if not output_tokens and not target_tokens:
212
+ return Score(name=self.name, value=0.0, kind=self.kind)
213
+ out_set = set(output_tokens)
214
+ tgt_set = set(target_tokens)
215
+ intersection = out_set & tgt_set
216
+ precision = len(intersection) / len(output_tokens) if output_tokens else 0.0
217
+ recall = len(intersection) / len(target_tokens) if target_tokens else 0.0
218
+ if precision == 0.0 and recall == 0.0:
219
+ f1 = 0.0
220
+ else:
221
+ f1 = 2.0 * precision * recall / (precision + recall)
222
+ return Score(name=self.name, value=f1, kind=self.kind)
@@ -0,0 +1,131 @@
1
+ """Embedding-based similarity metrics."""
2
+ from __future__ import annotations
3
+
4
+ from typing import Any
5
+
6
+ from auditkit.metric import Metric
7
+ from auditkit.registry import METRICS
8
+ from auditkit.sample import Sample
9
+ from auditkit.score import Score
10
+ from auditkit.types import Direction, ScoreKind
11
+ from auditkit.errors import ExtraNotInstalled
12
+
13
+
14
+ @METRICS.register("cosine_similarity")
15
+ class CosineSimilarity(Metric):
16
+ """Embedding cosine similarity via plain ``transformers`` -- no
17
+ ``sentence-transformers`` dependency.
18
+
19
+ A previous version of this metric used the ``sentence-transformers``
20
+ package, which transitively imports ``transformers``' audio/video
21
+ processing modules and, in turn, ``torchcodec`` -- a real, live-
22
+ confirmed environment failure on Colab (``torch``/``torchcodec``
23
+ version mismatch, unrelated ``libavutil``/FFmpeg shared libraries
24
+ missing) that has nothing to do with text embeddings at all and took
25
+ this metric down with it. ``sentence-transformers`` was used for
26
+ exactly this one metric in the whole library (confirmed: no other
27
+ file imports it), so removing it entirely and reimplementing the
28
+ same computation directly on ``transformers.AutoModel`` -- already a
29
+ core dependency for :class:`EncoderJudge`/``HFGenModel``, and
30
+ confirmed *not* to trigger the ``torchcodec`` import path -- avoids
31
+ the whole problem rather than working around it.
32
+
33
+ Mean-pooling + L2-normalization is the exact recipe
34
+ ``sentence-transformers`` itself uses internally for MiniLM-family
35
+ checkpoints (documented on the model card) -- this reproduces the
36
+ same embeddings, not an approximation.
37
+ """
38
+
39
+ name = "cosine_similarity"
40
+ kind = ScoreKind.BENCHMARK
41
+ direction = Direction.MAXIMIZE
42
+ is_deterministic = False
43
+ required_fields = frozenset({"target"})
44
+
45
+ def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2") -> None:
46
+ self._model_name = model_name
47
+ self._model = None
48
+ self._tokenizer = None
49
+ self._torch = None
50
+
51
+ def identity(self) -> dict:
52
+ return {"name": self.name, "model_name": self._model_name}
53
+
54
+ def _ensure_model(self) -> None:
55
+ if self._model is not None:
56
+ return
57
+ try:
58
+ import torch
59
+ from transformers import AutoModel, AutoTokenizer
60
+ except ImportError:
61
+ raise ExtraNotInstalled("transformers", "pip install auditkit[transformers]")
62
+ self._torch = torch
63
+ self._tokenizer = AutoTokenizer.from_pretrained(self._model_name)
64
+ self._model = AutoModel.from_pretrained(self._model_name)
65
+ self._model.eval()
66
+
67
+ def _embed(self, text: str) -> Any:
68
+ torch = self._torch
69
+ inputs = self._tokenizer([text], padding=True, truncation=True, return_tensors="pt")
70
+ with torch.no_grad():
71
+ output = self._model(**inputs)
72
+ # Mean pooling over real (non-padding) token embeddings -- the
73
+ # attention mask zeroes out padding positions before averaging,
74
+ # exactly as sentence-transformers does for this checkpoint family.
75
+ token_embeddings = output.last_hidden_state
76
+ mask = inputs["attention_mask"].unsqueeze(-1).expand(token_embeddings.size()).float()
77
+ summed = (token_embeddings * mask).sum(dim=1)
78
+ counts = mask.sum(dim=1).clamp(min=1e-9)
79
+ embedding = (summed / counts)[0]
80
+ norm = embedding.norm(p=2)
81
+ if norm > 0:
82
+ embedding = embedding / norm
83
+ return embedding
84
+
85
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
86
+ self._ensure_model()
87
+ emb_out = self._embed(output)
88
+ emb_tgt = self._embed(sample.target or "")
89
+ similarity = float(self._torch.dot(emb_out, emb_tgt).item())
90
+ return Score(name=self.name, value=similarity, kind=self.kind)
91
+
92
+
93
+ @METRICS.register("token_overlap")
94
+ class TokenOverlap(Metric):
95
+ name = "token_overlap"
96
+ kind = ScoreKind.BENCHMARK
97
+ direction = Direction.MAXIMIZE
98
+ is_deterministic = True
99
+ required_fields = frozenset({"target"})
100
+
101
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
102
+ tokens_out = set(output.split())
103
+ tokens_tgt = set((sample.target or "").split())
104
+ if not tokens_out and not tokens_tgt:
105
+ return Score(name=self.name, value=0.0, kind=self.kind)
106
+ intersection = tokens_out & tokens_tgt
107
+ union = tokens_out | tokens_tgt
108
+ return Score(name=self.name, value=len(intersection) / len(union), kind=self.kind)
109
+
110
+
111
+ @METRICS.register("bm25_similarity")
112
+ class BM25Similarity(Metric):
113
+ name = "bm25_similarity"
114
+ kind = ScoreKind.BENCHMARK
115
+ direction = Direction.MAXIMIZE
116
+ is_deterministic = True
117
+ required_fields = frozenset({"target"})
118
+
119
+ def score(self, sample: Sample, output: str, context: Any = None) -> Score:
120
+ tokens_out = output.split()
121
+ tokens_tgt = (sample.target or "").split()
122
+ freq_out: dict[str, int] = {}
123
+ for t in tokens_out:
124
+ freq_out[t] = freq_out.get(t, 0) + 1
125
+ freq_tgt: dict[str, int] = {}
126
+ for t in tokens_tgt:
127
+ freq_tgt[t] = freq_tgt.get(t, 0) + 1
128
+ all_tokens = set(freq_out) | set(freq_tgt)
129
+ score_val = sum(min(freq_out.get(t, 0), freq_tgt.get(t, 0)) for t in all_tokens)
130
+ max_len = max(len(tokens_out), len(tokens_tgt), 1)
131
+ return Score(name=self.name, value=score_val / max_len, kind=self.kind)