evalmetry 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.
evalmetry/judges.py ADDED
@@ -0,0 +1,273 @@
1
+ """Explicit, per-dataset LLM grading after target-model generation is saved.
2
+
3
+ The transport accepts a chat-completions HTTP endpoint. Dataset functions own
4
+ the messages and score interpretation; the framework owns persistence, retry,
5
+ metric validation and harness aggregation. No judge is selected automatically.
6
+ """
7
+ from __future__ import annotations
8
+
9
+ import hashlib
10
+ import json
11
+ import math
12
+ import os
13
+ from collections import defaultdict
14
+ from datetime import datetime, timezone
15
+ from http.client import HTTPException
16
+ from pathlib import Path
17
+ import time
18
+ from urllib.error import HTTPError, URLError
19
+ from urllib.parse import urlsplit
20
+ from urllib.request import Request, urlopen
21
+
22
+
23
+ def validate_judge_config(task: dict, protocol: dict) -> None:
24
+ """Reject ambiguous modes and invalid transport settings before model loading."""
25
+ mode = protocol.get('scoring', 'standard')
26
+ if mode not in ('standard', 'llm_judge'):
27
+ raise ValueError('scoring must be standard or llm_judge')
28
+ if mode == 'standard':
29
+ if 'judge' in protocol:
30
+ raise ValueError('judge settings require scoring: llm_judge')
31
+ return
32
+ if task.get('output_type') != 'generate_until':
33
+ raise ValueError('llm_judge requires generate_until (grade generated answers)')
34
+ if 'process_results' in task:
35
+ raise ValueError('llm_judge uses judge.score; remove top-level process_results')
36
+ judge = protocol.get('judge')
37
+ if not isinstance(judge, dict):
38
+ raise ValueError('llm_judge requires a judge mapping')
39
+ allowed = {'endpoint', 'model', 'api_key_env', 'prompt', 'score',
40
+ 'generation_kwargs', 'timeout', 'max_retries'}
41
+ if set(judge) - allowed:
42
+ raise ValueError(f'unknown judge settings: {sorted(set(judge) - allowed)}')
43
+ for key in ('endpoint', 'model', 'prompt', 'score'):
44
+ if not isinstance(judge.get(key), str) or not judge[key].strip():
45
+ raise ValueError(f'judge.{key} is required')
46
+ url = urlsplit(judge['endpoint'])
47
+ if url.scheme not in ('http', 'https') or not url.hostname or url.username or url.password or url.query or url.fragment:
48
+ raise ValueError('judge.endpoint must be an HTTP(S) URL without credentials/query/fragment')
49
+ key_env = judge.get('api_key_env')
50
+ if key_env is not None and (not isinstance(key_env, str) or not key_env.isidentifier()):
51
+ raise ValueError('judge.api_key_env must name an environment variable')
52
+ timeout = judge.get('timeout', 60)
53
+ if type(timeout) not in (int, float) or not math.isfinite(timeout) or timeout <= 0:
54
+ raise ValueError('judge.timeout must be positive and finite')
55
+ retries = judge.get('max_retries', 2)
56
+ if type(retries) is not int or not 0 <= retries <= 10:
57
+ raise ValueError('judge.max_retries must be an integer between 0 and 10')
58
+ kwargs = judge.get('generation_kwargs', {})
59
+ if not isinstance(kwargs, dict) or {'model', 'messages', 'stream'}.intersection(kwargs):
60
+ raise ValueError('judge.generation_kwargs cannot override model/messages/stream')
61
+ json.dumps(kwargs, allow_nan=False)
62
+ metrics = task.get('metric_list', [])
63
+ if any(not isinstance(m.get('metric'), str) for m in metrics):
64
+ raise ValueError('judge metrics must have string names')
65
+
66
+
67
+ def pending_judge_results(doc, results):
68
+ """Return no provisional scores: ungraded answers must never appear as zero."""
69
+ return {}
70
+
71
+
72
+ def _write_json(path: Path, value: dict) -> None:
73
+ """Replace a checkpoint atomically, keeping a complete file after interruption."""
74
+ path.parent.mkdir(parents=True, exist_ok=True)
75
+ temporary = path.with_suffix('.tmp')
76
+ temporary.write_text(json.dumps(value, ensure_ascii=False, allow_nan=False), encoding='utf-8')
77
+ temporary.replace(path)
78
+
79
+
80
+ def _reject_nonfinite(token: str):
81
+ raise ValueError(f'judge HTTP response contains non-finite JSON number {token}')
82
+
83
+
84
+ def _finite_float(text: str) -> float:
85
+ value = float(text)
86
+ if not math.isfinite(value):
87
+ _reject_nonfinite(text)
88
+ return value
89
+
90
+
91
+ def _safe_http_error_detail(error: HTTPError, payload: dict) -> str:
92
+ """Expose only status and a request-key name, never the server's message."""
93
+ detail = f"HTTP {error.code}"
94
+ try:
95
+ body = json.load(error, parse_constant=_reject_nonfinite, parse_float=_finite_float)
96
+ parameter = body.get('error', {}).get('param') if isinstance(body, dict) else None
97
+ except Exception:
98
+ parameter = None
99
+ # The response may be hostile or echo secrets. A key already present in our
100
+ # payload is safe to name; arbitrary response strings are not.
101
+ if isinstance(parameter, str) and parameter in payload:
102
+ detail += f", parameter {parameter!r}"
103
+ return detail
104
+
105
+
106
+ def request_judge(config: dict, messages: list) -> dict:
107
+ """Call an explicitly configured endpoint; only transient transport errors retry.
108
+
109
+ API keys are read at request time and never included in saved configuration.
110
+ Invalid model output is preserved for diagnosis instead of sampling repeatedly
111
+ until a parseable (and potentially biased) grade appears.
112
+ """
113
+ headers = {'Content-Type': 'application/json'}
114
+ if config.get('api_key_env'):
115
+ key = os.environ.get(config['api_key_env'])
116
+ if not key:
117
+ raise ValueError(f"missing judge API key environment variable: {config['api_key_env']}")
118
+ headers['Authorization'] = f'Bearer {key}'
119
+ payload = {'model': config['model'], 'messages': messages, 'stream': False,
120
+ **config.get('generation_kwargs', {})}
121
+ request = Request(config['endpoint'], data=json.dumps(payload, allow_nan=False).encode(),
122
+ headers=headers, method='POST')
123
+ retries = config.get('max_retries', 2)
124
+ for attempt in range(retries + 1):
125
+ try:
126
+ with urlopen(request, timeout=config.get('timeout', 60)) as response:
127
+ # NaN/Infinity would later make the failure record unwritable.
128
+ result = json.load(response, parse_constant=_reject_nonfinite,
129
+ parse_float=_finite_float)
130
+ if not isinstance(result, dict):
131
+ raise ValueError('judge HTTP response must be a JSON object')
132
+ return result
133
+ # urllib raises a dropped connection or truncated body outside URLError.
134
+ except (HTTPError, URLError, TimeoutError, ConnectionError, HTTPException) as error:
135
+ transient = not isinstance(error, HTTPError) or error.code in (408, 429, 500, 502, 503, 504)
136
+ if not transient or attempt == retries:
137
+ # Do not expose server messages, which could echo credentials.
138
+ detail = _safe_http_error_detail(error, payload) if isinstance(error, HTTPError) else None
139
+ suffix = f" ({detail})" if detail else ""
140
+ raise RuntimeError(f'judge request failed{suffix}; no score was assigned') from None
141
+ time.sleep(min(2 ** attempt, 8))
142
+ raise AssertionError('unreachable')
143
+
144
+
145
+ def _validate_scores(scores, metric_names: set, correctness: str | None) -> None:
146
+ """Require the declared finite scalar metrics, with explicit binary correctness."""
147
+ if not isinstance(scores, dict) or set(scores) != metric_names:
148
+ raise ValueError('judge.score must return exactly the metric_list names')
149
+ for name, value in scores.items():
150
+ if not isinstance(value, (int, float)) or not math.isfinite(value):
151
+ raise ValueError(f'judge metric {name} must be a finite number')
152
+ if correctness and scores[correctness] not in (0, 1):
153
+ raise ValueError('judge correctness_metric must be binary 0/1')
154
+
155
+
156
+ def grade_sample(sample: dict, protocol: dict, metric_names: set,
157
+ run_dir: Path, provenance: dict) -> dict:
158
+ """Grade one filtered answer, reusing only a successful identical judge request.
159
+
160
+ Full document, filtered answers, rendered messages, protocol and source hashes
161
+ participate in cache identity. A changed rubric, parser or reference answer
162
+ therefore cannot silently reuse an old score. The response is saved before
163
+ calling user score code, so parsing failures remain inspectable.
164
+ """
165
+ judge = protocol['judge']
166
+ messages = judge['prompt'](sample['doc'], sample['filtered_resps'])
167
+ if not isinstance(messages, list) or not messages or any(
168
+ not isinstance(m, dict) or m.get('role') not in ('system', 'user', 'assistant')
169
+ or not isinstance(m.get('content'), str) for m in messages
170
+ ):
171
+ raise ValueError('judge.prompt must return nonempty chat messages with role/content')
172
+ transport = {key: value for key, value in judge.items() if key not in ('prompt', 'score')}
173
+ identity = {'provenance': provenance, 'judge': transport, 'task': sample.get('task_name'),
174
+ 'doc': sample['doc'], 'responses': sample['filtered_resps'],
175
+ 'filter': sample.get('filter', 'none'), 'messages': messages,
176
+ 'metrics': sorted(metric_names)}
177
+ digest = hashlib.sha256(json.dumps(identity, sort_keys=True, ensure_ascii=False,
178
+ allow_nan=False).encode()).hexdigest()
179
+ path = run_dir / 'judge' / f'{digest}.json'
180
+ if path.exists():
181
+ record = json.loads(path.read_text())
182
+ if record.get('status') == 'success':
183
+ _validate_scores(record['scores'], metric_names, protocol.get('correctness_metric'))
184
+ sample['_judge'] = {'status': 'success', 'record': str(path.relative_to(run_dir)),
185
+ 'cache_hit': True}
186
+ return record['scores']
187
+ # Keep the service call time outside the cache identity: rerunning an identical
188
+ # request may reuse this record, while the timestamp documents the original call.
189
+ record = {'status': 'pending', 'identity': identity,
190
+ 'requested_at': datetime.now(timezone.utc).isoformat()}
191
+ _write_json(path, record)
192
+ try:
193
+ record['response'] = request_judge(judge, messages)
194
+ _write_json(path, record)
195
+ scores = judge['score'](sample['doc'], sample['filtered_resps'], record['response'])
196
+ _validate_scores(scores, metric_names, protocol.get('correctness_metric'))
197
+ record.update(status='success', scores=scores)
198
+ except Exception as error:
199
+ record.update(status='failed', error_type=type(error).__name__)
200
+ _write_json(path, record)
201
+ sample['_judge'] = {'status': 'failed', 'record': str(path.relative_to(run_dir))}
202
+ raise
203
+ _write_json(path, record)
204
+ sample['_judge'] = {'status': 'success', 'record': str(path.relative_to(run_dir)),
205
+ 'cache_hit': False}
206
+ return scores
207
+
208
+
209
+ def score_judge_tasks(results: dict, tasks: list, run_dir, provenance: dict) -> None:
210
+ """Persist all generations first, then grade and aggregate through lm-eval.
211
+
212
+ On failure the run stays incomplete; successful per-sample scores remain in
213
+ the journal and no partial aggregate is published. Ordinary tasks are untouched.
214
+ """
215
+ from . import storage
216
+
217
+ judges = [task for task in tasks if not isinstance(task, (str, dict))
218
+ and (task.config.metadata or {}).get('eval_framework', {}).get('scoring') == 'llm_judge']
219
+ if not judges:
220
+ return
221
+ samples_by_task = results.get('samples', {})
222
+ for task in judges:
223
+ protocol = task.config.metadata['eval_framework']
224
+ for sample in samples_by_task.get(task.config.task, []):
225
+ sample['_judge'] = {'status': 'pending'}
226
+ sample['_eval_framework'] = {
227
+ 'sample_id': sample['doc'][protocol['sample_id']],
228
+ 'primary_filter': protocol['primary_filter'], 'is_correct': None,
229
+ }
230
+ for metric in task.aggregation():
231
+ sample.pop(metric, None)
232
+ sample['metrics'] = []
233
+ storage.write_samples(run_dir, samples_by_task)
234
+ try:
235
+ for task in judges:
236
+ name = task.config.task
237
+ protocol = task.config.metadata['eval_framework']
238
+ sample_metrics = defaultdict(list)
239
+ for sample in samples_by_task.get(name, []):
240
+ grading_sample = {**sample, 'task_name': name}
241
+ try:
242
+ scores = grade_sample(grading_sample, protocol, set(task.aggregation()),
243
+ Path(run_dir), provenance)
244
+ except Exception:
245
+ grading_sample['_judge']['status'] = 'failed'
246
+ raise
247
+ finally:
248
+ sample['_judge'] = grading_sample['_judge']
249
+ sample.update(scores)
250
+ sample['metrics'] = list(scores)
251
+ if protocol.get('correctness_metric'):
252
+ sample['_eval_framework']['is_correct'] = bool(scores[protocol['correctness_metric']])
253
+ for metric, value in scores.items():
254
+ sample_metrics[(metric, sample.get('filter', 'none'))].append(value)
255
+ # lm-eval 0.4.13 extracted this helper from the 0.4.9.1 TaskOutput.
256
+ # Use each installed version's own aggregation and stderr behavior.
257
+ try:
258
+ from lm_eval.evaluator_utils import _compute_task_aggregations
259
+ except ImportError:
260
+ from lm_eval.evaluator_utils import TaskOutput
261
+ output = TaskOutput.from_taskdict(name, task)
262
+ output.sample_metrics = sample_metrics
263
+ output.calculate_aggregate_metric()
264
+ aggregate, count = output.agg_metrics, output.sample_len
265
+ count_key = 'samples'
266
+ else:
267
+ aggregate, count = _compute_task_aggregations(task, sample_metrics, 100000)
268
+ count_key = 'sample_len'
269
+ results['results'].setdefault(name, {}).update({**aggregate, count_key: count})
270
+ if name in results.get('n-samples', {}):
271
+ results['n-samples'][name]['effective'] = count
272
+ finally:
273
+ storage.write_samples(run_dir, samples_by_task)