synthbench-eval 0.4.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.
- synthbench/__init__.py +16 -0
- synthbench/__main__.py +5 -0
- synthbench/adapter.py +131 -0
- synthbench/anomaly.py +503 -0
- synthbench/baseline_floors.py +213 -0
- synthbench/baselines.py +902 -0
- synthbench/cli.py +2929 -0
- synthbench/config_id.py +424 -0
- synthbench/contamination.py +783 -0
- synthbench/convergence/__init__.py +67 -0
- synthbench/convergence/baseline.py +172 -0
- synthbench/convergence/bootstrap.py +60 -0
- synthbench/convergence/cli_report.py +556 -0
- synthbench/convergence/curves.py +111 -0
- synthbench/convergence/real_sampling.py +146 -0
- synthbench/convergence/thresholds.py +49 -0
- synthbench/datasets/__init__.py +37 -0
- synthbench/datasets/base.py +147 -0
- synthbench/datasets/eurobarometer.py +335 -0
- synthbench/datasets/globalopinionqa.py +252 -0
- synthbench/datasets/gss.py +323 -0
- synthbench/datasets/michigan.py +505 -0
- synthbench/datasets/ntia.py +357 -0
- synthbench/datasets/opinionsqa.py +412 -0
- synthbench/datasets/pewtech.py +334 -0
- synthbench/datasets/policy.py +142 -0
- synthbench/datasets/subpop.py +302 -0
- synthbench/datasets/wvs.py +229 -0
- synthbench/findings.py +988 -0
- synthbench/holdout.py +347 -0
- synthbench/human_distributions.py +243 -0
- synthbench/leaderboard.py +715 -0
- synthbench/leaderboard_pr.py +274 -0
- synthbench/metrics/__init__.py +38 -0
- synthbench/metrics/composite.py +72 -0
- synthbench/metrics/conditioning.py +39 -0
- synthbench/metrics/distributional.py +41 -0
- synthbench/metrics/ranking.py +37 -0
- synthbench/metrics/refusal.py +270 -0
- synthbench/metrics/subgroup.py +58 -0
- synthbench/private_holdout.py +240 -0
- synthbench/providers/__init__.py +43 -0
- synthbench/providers/_parsing.py +150 -0
- synthbench/providers/_retry.py +107 -0
- synthbench/providers/base.py +212 -0
- synthbench/providers/http.py +108 -0
- synthbench/providers/majority_baseline.py +23 -0
- synthbench/providers/ollama.py +109 -0
- synthbench/providers/openrouter.py +144 -0
- synthbench/providers/population_baseline.py +87 -0
- synthbench/providers/random_baseline.py +24 -0
- synthbench/providers/raw_anthropic.py +149 -0
- synthbench/providers/raw_gemini.py +145 -0
- synthbench/providers/raw_openai.py +139 -0
- synthbench/providers/synthpanel.py +978 -0
- synthbench/publish.py +2772 -0
- synthbench/r2_upload.py +178 -0
- synthbench/recompute.py +235 -0
- synthbench/report.py +493 -0
- synthbench/run_hash.py +111 -0
- synthbench/run_validity.py +232 -0
- synthbench/runner.py +834 -0
- synthbench/stats.py +1492 -0
- synthbench/submission.py +284 -0
- synthbench/submission_pr.py +436 -0
- synthbench/submit_adapter.py +748 -0
- synthbench/suite.py +424 -0
- synthbench/suites/__init__.py +94 -0
- synthbench/topics.py +256 -0
- synthbench/user_config.py +409 -0
- synthbench/validation.py +1556 -0
- synthbench/visualize.py +455 -0
- synthbench_eval-0.4.0.dist-info/METADATA +279 -0
- synthbench_eval-0.4.0.dist-info/RECORD +78 -0
- synthbench_eval-0.4.0.dist-info/WHEEL +5 -0
- synthbench_eval-0.4.0.dist-info/entry_points.txt +2 -0
- synthbench_eval-0.4.0.dist-info/licenses/LICENSE +21 -0
- synthbench_eval-0.4.0.dist-info/top_level.txt +1 -0
synthbench/baselines.py
ADDED
|
@@ -0,0 +1,902 @@
|
|
|
1
|
+
"""Human Ceiling, Temporal Drift Floor, and Ensemble Bootstrap CI baselines.
|
|
2
|
+
|
|
3
|
+
Implements the split-half multinomial bootstrap ceiling described in the
|
|
4
|
+
methodology writeup (hq-wisp-vuom1):
|
|
5
|
+
|
|
6
|
+
1. compute_ceiling() — pure function. Given raw category counts, draws two
|
|
7
|
+
independent half-samples from Multinomial(n/2, p_hat), applies a distance
|
|
8
|
+
metric between them, and reports (1 - mean(distance)) as the ceiling.
|
|
9
|
+
|
|
10
|
+
2. compute_temporal_drift() — JSD between same-wording questions across Pew
|
|
11
|
+
ATP waves. Reports mean drift per year-gap + CI. A new baseline-adjacent
|
|
12
|
+
metric that quantifies how real-world opinions shift year-over-year.
|
|
13
|
+
|
|
14
|
+
3. ensemble_bootstrap_ci() — resamples per-question scores with replacement
|
|
15
|
+
B=1000 times to produce percentile CIs for deterministic ensemble runs
|
|
16
|
+
(previously reported CI_lower = CI_upper = 0.000).
|
|
17
|
+
|
|
18
|
+
Citations (per data scientist):
|
|
19
|
+
- Santurkar et al. 2023 (OpinionsQA, arxiv:2303.17548)
|
|
20
|
+
- Durmus et al. 2023 (GlobalOpinionQA, arxiv:2306.16388)
|
|
21
|
+
- Suh et al. (SubPOP, ACL 2025, arxiv:2502.16761)
|
|
22
|
+
- Spearman-Brown prophecy (1910)
|
|
23
|
+
- Efron (1979) bootstrap
|
|
24
|
+
- Lin (1991) Jensen-Shannon divergence
|
|
25
|
+
- Cochran (1977) Sampling Techniques
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
from __future__ import annotations
|
|
29
|
+
|
|
30
|
+
from dataclasses import dataclass
|
|
31
|
+
from typing import Callable, Literal
|
|
32
|
+
|
|
33
|
+
import numpy as np
|
|
34
|
+
from scipy.special import rel_entr
|
|
35
|
+
|
|
36
|
+
from synthbench.metrics.distributional import jensen_shannon_divergence
|
|
37
|
+
from synthbench.metrics.ranking import kendall_tau_b
|
|
38
|
+
|
|
39
|
+
_LOG2 = float(np.log(2.0))
|
|
40
|
+
|
|
41
|
+
QualityFlag = Literal["high", "medium", "low"]
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@dataclass
|
|
45
|
+
class CeilingResult:
|
|
46
|
+
"""Result of a split-half ceiling computation."""
|
|
47
|
+
|
|
48
|
+
mean: float
|
|
49
|
+
ci_low: float
|
|
50
|
+
ci_high: float
|
|
51
|
+
n_effective: int
|
|
52
|
+
quality_flag: QualityFlag
|
|
53
|
+
method: str = "multinomial_bootstrap_1000"
|
|
54
|
+
|
|
55
|
+
def to_dict(self) -> dict:
|
|
56
|
+
return {
|
|
57
|
+
"mean": round(self.mean, 6),
|
|
58
|
+
"ci_low": round(self.ci_low, 6),
|
|
59
|
+
"ci_high": round(self.ci_high, 6),
|
|
60
|
+
"n_effective": self.n_effective,
|
|
61
|
+
"quality_flag": self.quality_flag,
|
|
62
|
+
"method": self.method,
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _quality_flag(n: int) -> QualityFlag:
|
|
67
|
+
"""Classify raw subpop sample size per Cochran (1977) and methodology writeup."""
|
|
68
|
+
if n >= 400:
|
|
69
|
+
return "high"
|
|
70
|
+
if n >= 200:
|
|
71
|
+
return "medium"
|
|
72
|
+
return "low"
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def compute_ceiling(
|
|
76
|
+
counts: dict[str, int],
|
|
77
|
+
metric: Callable[[dict[str, float], dict[str, float]], float],
|
|
78
|
+
n_bootstrap: int = 1000,
|
|
79
|
+
seed: int = 42,
|
|
80
|
+
is_distance: bool = True,
|
|
81
|
+
) -> CeilingResult:
|
|
82
|
+
"""Split-half ceiling via multinomial bootstrap.
|
|
83
|
+
|
|
84
|
+
Given observed counts c = [c_1, ..., c_k] with total n, treats p_hat = c/n
|
|
85
|
+
as the multinomial MLE. Draws two independent half-samples of size
|
|
86
|
+
floor(n/2) from Multinomial(n/2, p_hat), computes the metric between the
|
|
87
|
+
two empirical distributions, and repeats B times.
|
|
88
|
+
|
|
89
|
+
For distance metrics (JSD, where lower = better), ceiling = 1 - mean(d).
|
|
90
|
+
For agreement metrics (Kendall tau, where higher = better), ceiling =
|
|
91
|
+
mean(tau) directly — pass is_distance=False.
|
|
92
|
+
|
|
93
|
+
Args:
|
|
94
|
+
counts: option -> count from real human data. Integer-valued.
|
|
95
|
+
metric: callable(p, q) -> float, where p/q are dict[option, probability].
|
|
96
|
+
n_bootstrap: number of bootstrap replicates (default 1000).
|
|
97
|
+
seed: RNG seed for reproducibility (default 42).
|
|
98
|
+
is_distance: If True, ceiling = 1 - mean(metric). If False, ceiling =
|
|
99
|
+
mean(metric) (useful for rank agreement like Kendall tau).
|
|
100
|
+
|
|
101
|
+
Returns:
|
|
102
|
+
CeilingResult with mean, 95% CI, effective n, quality flag, method.
|
|
103
|
+
|
|
104
|
+
Raises:
|
|
105
|
+
ValueError: if counts is empty or all zero.
|
|
106
|
+
"""
|
|
107
|
+
keys = sorted(counts.keys())
|
|
108
|
+
if not keys:
|
|
109
|
+
raise ValueError("counts is empty")
|
|
110
|
+
|
|
111
|
+
c_vec = np.array([float(counts[k]) for k in keys], dtype=np.float64)
|
|
112
|
+
n_total = int(c_vec.sum())
|
|
113
|
+
if n_total <= 0:
|
|
114
|
+
raise ValueError(f"counts sum to {n_total}, must be positive")
|
|
115
|
+
|
|
116
|
+
p_hat = c_vec / n_total
|
|
117
|
+
n_half = n_total // 2
|
|
118
|
+
|
|
119
|
+
# Degenerate case: n too small to split
|
|
120
|
+
if n_half < 1:
|
|
121
|
+
return CeilingResult(
|
|
122
|
+
mean=0.0,
|
|
123
|
+
ci_low=0.0,
|
|
124
|
+
ci_high=0.0,
|
|
125
|
+
n_effective=n_total,
|
|
126
|
+
quality_flag=_quality_flag(n_total),
|
|
127
|
+
method=f"multinomial_bootstrap_{n_bootstrap}",
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
rng = np.random.default_rng(seed)
|
|
131
|
+
|
|
132
|
+
if metric is jensen_shannon_divergence:
|
|
133
|
+
# Fast path: vectorize both the multinomial draws and the JSD metric.
|
|
134
|
+
# `size=(n_bootstrap, 2)` preserves the interleaved A/B draw order of
|
|
135
|
+
# the original per-iteration loop, so the RNG bit-stream — and thus
|
|
136
|
+
# every realized draw — is identical to calling rng.multinomial twice
|
|
137
|
+
# per step. This keeps seed=42 reproducibility bit-exact vs. the
|
|
138
|
+
# dict-based path (see sb-dkz).
|
|
139
|
+
distances = _multinomial_bootstrap_jsd(p_hat, n_half, n_bootstrap, rng)
|
|
140
|
+
else:
|
|
141
|
+
distances = np.empty(n_bootstrap, dtype=np.float64)
|
|
142
|
+
for b in range(n_bootstrap):
|
|
143
|
+
# Two independent half-samples from Multinomial(n/2, p_hat)
|
|
144
|
+
draw_a = rng.multinomial(n_half, p_hat)
|
|
145
|
+
draw_b = rng.multinomial(n_half, p_hat)
|
|
146
|
+
|
|
147
|
+
# Normalize to probabilities for the metric
|
|
148
|
+
dist_a = {keys[i]: float(draw_a[i]) / n_half for i in range(len(keys))}
|
|
149
|
+
dist_b = {keys[i]: float(draw_b[i]) / n_half for i in range(len(keys))}
|
|
150
|
+
|
|
151
|
+
distances[b] = float(metric(dist_a, dist_b))
|
|
152
|
+
|
|
153
|
+
mean_d = float(distances.mean())
|
|
154
|
+
lo_arr, hi_arr = np.percentile(distances, [2.5, 97.5])
|
|
155
|
+
lo = float(lo_arr)
|
|
156
|
+
hi = float(hi_arr)
|
|
157
|
+
|
|
158
|
+
if is_distance:
|
|
159
|
+
# Ceiling = 1 - mean(distance); CI flips direction
|
|
160
|
+
ceiling_mean = 1.0 - mean_d
|
|
161
|
+
ceiling_lo = 1.0 - hi
|
|
162
|
+
ceiling_hi = 1.0 - lo
|
|
163
|
+
else:
|
|
164
|
+
# Agreement metric: ceiling = mean(metric) directly
|
|
165
|
+
ceiling_mean = mean_d
|
|
166
|
+
ceiling_lo = lo
|
|
167
|
+
ceiling_hi = hi
|
|
168
|
+
|
|
169
|
+
return CeilingResult(
|
|
170
|
+
mean=ceiling_mean,
|
|
171
|
+
ci_low=ceiling_lo,
|
|
172
|
+
ci_high=ceiling_hi,
|
|
173
|
+
n_effective=n_total,
|
|
174
|
+
quality_flag=_quality_flag(n_total),
|
|
175
|
+
method=f"multinomial_bootstrap_{n_bootstrap}",
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def _multinomial_bootstrap_jsd(
|
|
180
|
+
p_hat: np.ndarray,
|
|
181
|
+
n_half: int,
|
|
182
|
+
n_bootstrap: int,
|
|
183
|
+
rng: np.random.Generator,
|
|
184
|
+
) -> np.ndarray:
|
|
185
|
+
"""Vectorized split-half JSD bootstrap (base-2), bit-compatible with
|
|
186
|
+
the per-iteration loop using two sequential ``rng.multinomial(n_half,
|
|
187
|
+
p_hat)`` calls at the same seed.
|
|
188
|
+
|
|
189
|
+
``size=(n_bootstrap, 2)`` materializes the draws in the same interleaved
|
|
190
|
+
A, B, A, B, ... order the original loop consumed, so the RNG state is
|
|
191
|
+
advanced identically. The returned JSD values match the dict-based
|
|
192
|
+
``jensen_shannon_divergence(p, q)`` computation to ~1e-16 (floating-point
|
|
193
|
+
noise only), which is well under the 6-decimal rounding applied in
|
|
194
|
+
``CeilingResult.to_dict``.
|
|
195
|
+
"""
|
|
196
|
+
draws = rng.multinomial(n_half, p_hat, size=(n_bootstrap, 2))
|
|
197
|
+
p = draws[:, 0].astype(np.float64) / n_half
|
|
198
|
+
q = draws[:, 1].astype(np.float64) / n_half
|
|
199
|
+
m = 0.5 * (p + q)
|
|
200
|
+
jsd_nats = 0.5 * (rel_entr(p, m).sum(axis=-1) + rel_entr(q, m).sum(axis=-1))
|
|
201
|
+
return jsd_nats / _LOG2
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
def compute_ceiling_jsd(
|
|
205
|
+
counts: dict[str, int], n_bootstrap: int = 1000, seed: int = 42
|
|
206
|
+
) -> CeilingResult:
|
|
207
|
+
"""Convenience wrapper for JSD-based ceiling (P_dist)."""
|
|
208
|
+
return compute_ceiling(
|
|
209
|
+
counts, jensen_shannon_divergence, n_bootstrap, seed, is_distance=True
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
def compute_ceiling_tau(
|
|
214
|
+
counts: dict[str, int], n_bootstrap: int = 1000, seed: int = 42
|
|
215
|
+
) -> CeilingResult:
|
|
216
|
+
"""Convenience wrapper for Kendall tau-based ceiling (P_rank)."""
|
|
217
|
+
return compute_ceiling(counts, kendall_tau_b, n_bootstrap, seed, is_distance=False)
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
def aggregate_ceilings(
|
|
221
|
+
ceilings: list[CeilingResult],
|
|
222
|
+
weights: list[float] | None = None,
|
|
223
|
+
) -> CeilingResult | None:
|
|
224
|
+
"""Aggregate per-subpop/per-wave ceilings into a single summary.
|
|
225
|
+
|
|
226
|
+
Weighted mean of the point estimate; weighted CI bounds (keeps the range
|
|
227
|
+
interpretable without mis-implying independence across subpops). Method
|
|
228
|
+
string reflects aggregation.
|
|
229
|
+
|
|
230
|
+
Args:
|
|
231
|
+
ceilings: non-empty list of CeilingResult objects to combine.
|
|
232
|
+
weights: optional weights (e.g., n_questions per wave). If None,
|
|
233
|
+
weights by n_effective of each result.
|
|
234
|
+
|
|
235
|
+
Returns:
|
|
236
|
+
Aggregated CeilingResult, or None if input is empty.
|
|
237
|
+
"""
|
|
238
|
+
if not ceilings:
|
|
239
|
+
return None
|
|
240
|
+
|
|
241
|
+
if weights is None:
|
|
242
|
+
weights = [float(c.n_effective) for c in ceilings]
|
|
243
|
+
if len(weights) != len(ceilings):
|
|
244
|
+
raise ValueError("weights and ceilings must have same length")
|
|
245
|
+
|
|
246
|
+
w_sum = float(sum(weights))
|
|
247
|
+
if w_sum <= 0:
|
|
248
|
+
return None
|
|
249
|
+
|
|
250
|
+
mean = sum(c.mean * w for c, w in zip(ceilings, weights)) / w_sum
|
|
251
|
+
lo = sum(c.ci_low * w for c, w in zip(ceilings, weights)) / w_sum
|
|
252
|
+
hi = sum(c.ci_high * w for c, w in zip(ceilings, weights)) / w_sum
|
|
253
|
+
n_eff = int(sum(c.n_effective for c in ceilings))
|
|
254
|
+
|
|
255
|
+
# Aggregate quality flag = worst tier present
|
|
256
|
+
flags = {c.quality_flag for c in ceilings}
|
|
257
|
+
if "low" in flags:
|
|
258
|
+
agg_flag: QualityFlag = "low"
|
|
259
|
+
elif "medium" in flags:
|
|
260
|
+
agg_flag = "medium"
|
|
261
|
+
else:
|
|
262
|
+
agg_flag = "high"
|
|
263
|
+
|
|
264
|
+
return CeilingResult(
|
|
265
|
+
mean=mean,
|
|
266
|
+
ci_low=lo,
|
|
267
|
+
ci_high=hi,
|
|
268
|
+
n_effective=n_eff,
|
|
269
|
+
quality_flag=agg_flag,
|
|
270
|
+
method=f"aggregate_of_{len(ceilings)}_subpops",
|
|
271
|
+
)
|
|
272
|
+
|
|
273
|
+
|
|
274
|
+
def compute_temporal_drift(
|
|
275
|
+
per_question_data: list[dict],
|
|
276
|
+
) -> dict:
|
|
277
|
+
"""Temporal drift floor for OpinionsQA: JSD between same-wording questions
|
|
278
|
+
across Pew ATP waves.
|
|
279
|
+
|
|
280
|
+
Pew repeats ~15-20% of questions across waves for trend tracking. This
|
|
281
|
+
quantifies how much real opinions shift year-over-year on repeated items.
|
|
282
|
+
Useful framing for P_refuse and longitudinal claims; strictly separate
|
|
283
|
+
from the Human Ceiling itself.
|
|
284
|
+
|
|
285
|
+
Identifies repeated questions by stripping the wave suffix (e.g.
|
|
286
|
+
"TRUST_W32" -> "TRUST") and pairing distributions across waves.
|
|
287
|
+
|
|
288
|
+
Args:
|
|
289
|
+
per_question_data: list of dicts with at least keys
|
|
290
|
+
'key', 'human_distribution', and either 'temporal_year' (from
|
|
291
|
+
publish) or 'survey' (so the wave year can be extracted).
|
|
292
|
+
|
|
293
|
+
Returns:
|
|
294
|
+
dict with:
|
|
295
|
+
mean_drift: overall mean JSD across repeated-question pairs
|
|
296
|
+
ci_low, ci_high: 95% percentile CI via bootstrap over pairs
|
|
297
|
+
n_pairs: total number of cross-wave pairs compared
|
|
298
|
+
n_stems: number of question stems observed in 2+ waves
|
|
299
|
+
by_year_gap: dict mapping year-gap (int) -> mean JSD and count
|
|
300
|
+
method: description string
|
|
301
|
+
"""
|
|
302
|
+
from synthbench.datasets.opinionsqa import WAVE_YEAR_MAP, wave_year
|
|
303
|
+
|
|
304
|
+
# Group entries by key stem (strip "_W##" suffix). Pew sometimes assigns
|
|
305
|
+
# wave-specific keys to same-wording repeats — this groups them only when
|
|
306
|
+
# the underlying key (e.g., HARASS4) is reused across waves, which is
|
|
307
|
+
# Pew's convention for true trend-tracking questions. Text-based matching
|
|
308
|
+
# produces false positives because Pew questions often share generic
|
|
309
|
+
# preambles ("Please choose the statement that comes closer to your own
|
|
310
|
+
# views.") that aren't semantically equivalent.
|
|
311
|
+
stems: dict[str, list[dict]] = {}
|
|
312
|
+
for q in per_question_data:
|
|
313
|
+
key = q.get("key", "")
|
|
314
|
+
if not key or "_W" not in key:
|
|
315
|
+
continue
|
|
316
|
+
stem = key.rsplit("_W", 1)[0]
|
|
317
|
+
year = q.get("temporal_year")
|
|
318
|
+
if year is None or year == 0:
|
|
319
|
+
survey = q.get("survey", "")
|
|
320
|
+
year = wave_year(survey) if survey else 0
|
|
321
|
+
if not year:
|
|
322
|
+
try:
|
|
323
|
+
wave_num = int(key.rsplit("_W", 1)[1])
|
|
324
|
+
year = WAVE_YEAR_MAP.get(wave_num, 0)
|
|
325
|
+
except (ValueError, IndexError):
|
|
326
|
+
year = 0
|
|
327
|
+
if not year:
|
|
328
|
+
continue
|
|
329
|
+
dist = q.get("human_distribution", {})
|
|
330
|
+
if not dist:
|
|
331
|
+
continue
|
|
332
|
+
# Case-normalize stem: Pew sometimes uses SATLIFEA vs SATLIFEa for
|
|
333
|
+
# the same underlying trend item across waves.
|
|
334
|
+
stems.setdefault(stem.upper(), []).append({"year": int(year), "dist": dist})
|
|
335
|
+
|
|
336
|
+
# Compute JSD between every pair of waves for each repeated stem
|
|
337
|
+
pair_jsds: list[tuple[int, float]] = [] # (year_gap, jsd)
|
|
338
|
+
n_stems_repeated = 0
|
|
339
|
+
for stem, entries in stems.items():
|
|
340
|
+
if len(entries) < 2:
|
|
341
|
+
continue
|
|
342
|
+
n_stems_repeated += 1
|
|
343
|
+
for i in range(len(entries)):
|
|
344
|
+
for j in range(i + 1, len(entries)):
|
|
345
|
+
year_gap = abs(entries[i]["year"] - entries[j]["year"])
|
|
346
|
+
jsd = jensen_shannon_divergence(entries[i]["dist"], entries[j]["dist"])
|
|
347
|
+
pair_jsds.append((year_gap, jsd))
|
|
348
|
+
|
|
349
|
+
if not pair_jsds:
|
|
350
|
+
return {
|
|
351
|
+
"mean_drift": 0.0,
|
|
352
|
+
"ci_low": 0.0,
|
|
353
|
+
"ci_high": 0.0,
|
|
354
|
+
"n_pairs": 0,
|
|
355
|
+
"n_stems": 0,
|
|
356
|
+
"by_year_gap": {},
|
|
357
|
+
"method": "cross_wave_jsd_on_repeated_stems",
|
|
358
|
+
}
|
|
359
|
+
|
|
360
|
+
jsds = np.array([p[1] for p in pair_jsds], dtype=np.float64)
|
|
361
|
+
mean_drift = float(jsds.mean())
|
|
362
|
+
|
|
363
|
+
# Bootstrap CI over pairs
|
|
364
|
+
rng = np.random.default_rng(42)
|
|
365
|
+
n_boot = 1000
|
|
366
|
+
boot_means = np.empty(n_boot, dtype=np.float64)
|
|
367
|
+
n = len(jsds)
|
|
368
|
+
for b in range(n_boot):
|
|
369
|
+
idx = rng.integers(0, n, size=n)
|
|
370
|
+
boot_means[b] = float(jsds[idx].mean())
|
|
371
|
+
ci_low = float(np.percentile(boot_means, 2.5))
|
|
372
|
+
ci_high = float(np.percentile(boot_means, 97.5))
|
|
373
|
+
|
|
374
|
+
# By year-gap breakdown
|
|
375
|
+
by_gap: dict[int, list[float]] = {}
|
|
376
|
+
for gap, jsd in pair_jsds:
|
|
377
|
+
by_gap.setdefault(gap, []).append(jsd)
|
|
378
|
+
by_year_gap = {
|
|
379
|
+
str(gap): {
|
|
380
|
+
"mean_jsd": round(float(np.mean(vals)), 6),
|
|
381
|
+
"n_pairs": len(vals),
|
|
382
|
+
}
|
|
383
|
+
for gap, vals in sorted(by_gap.items())
|
|
384
|
+
}
|
|
385
|
+
|
|
386
|
+
return {
|
|
387
|
+
"mean_drift": round(mean_drift, 6),
|
|
388
|
+
"ci_low": round(ci_low, 6),
|
|
389
|
+
"ci_high": round(ci_high, 6),
|
|
390
|
+
"n_pairs": len(pair_jsds),
|
|
391
|
+
"n_stems": n_stems_repeated,
|
|
392
|
+
"by_year_gap": by_year_gap,
|
|
393
|
+
"method": "cross_wave_jsd_on_repeated_stems",
|
|
394
|
+
}
|
|
395
|
+
|
|
396
|
+
|
|
397
|
+
def ensemble_bootstrap_ci(
|
|
398
|
+
per_question: list[dict],
|
|
399
|
+
metric_key: str = "parity",
|
|
400
|
+
n_bootstrap: int = 1000,
|
|
401
|
+
seed: int = 42,
|
|
402
|
+
) -> tuple[float, float]:
|
|
403
|
+
"""Bootstrap CI for a deterministic ensemble run by resampling
|
|
404
|
+
per-question scores with replacement.
|
|
405
|
+
|
|
406
|
+
Ensemble entries currently emit CI_lower = CI_upper = 0.000 because the
|
|
407
|
+
arithmetic blend has no replicate variance. This function recovers a real
|
|
408
|
+
CI by bootstrapping over the per-question score distribution.
|
|
409
|
+
|
|
410
|
+
Args:
|
|
411
|
+
per_question: list of per-question dicts from an ensemble result.
|
|
412
|
+
metric_key: which per-question score to resample. Defaults to 'parity'
|
|
413
|
+
(matches the aggregate SPS / composite_parity).
|
|
414
|
+
n_bootstrap: number of bootstrap replicates (default 1000).
|
|
415
|
+
seed: RNG seed for reproducibility (default 42).
|
|
416
|
+
|
|
417
|
+
Returns:
|
|
418
|
+
(ci_lower, ci_upper) as 95% percentile CI. (0.0, 0.0) if no data.
|
|
419
|
+
"""
|
|
420
|
+
values = [
|
|
421
|
+
float(q[metric_key]) for q in per_question if q.get(metric_key) is not None
|
|
422
|
+
]
|
|
423
|
+
if not values:
|
|
424
|
+
return (0.0, 0.0)
|
|
425
|
+
|
|
426
|
+
arr = np.array(values, dtype=np.float64)
|
|
427
|
+
n = len(arr)
|
|
428
|
+
rng = np.random.default_rng(seed)
|
|
429
|
+
boot_means = np.empty(n_bootstrap, dtype=np.float64)
|
|
430
|
+
for b in range(n_bootstrap):
|
|
431
|
+
idx = rng.integers(0, n, size=n)
|
|
432
|
+
boot_means[b] = float(arr[idx].mean())
|
|
433
|
+
ci_low = float(np.percentile(boot_means, 2.5))
|
|
434
|
+
ci_high = float(np.percentile(boot_means, 97.5))
|
|
435
|
+
return (ci_low, ci_high)
|
|
436
|
+
|
|
437
|
+
|
|
438
|
+
# ---------------------------------------------------------------------------
|
|
439
|
+
# Per-dataset ceiling protocols
|
|
440
|
+
# ---------------------------------------------------------------------------
|
|
441
|
+
|
|
442
|
+
|
|
443
|
+
def _counts_from_probs(probs: dict[str, float], n: int) -> dict[str, int]:
|
|
444
|
+
"""Convert a probability distribution to integer counts for a given n.
|
|
445
|
+
|
|
446
|
+
Uses round-half-to-even; adjusts the largest bucket to guarantee the
|
|
447
|
+
counts sum to n (deterministic for fixed input).
|
|
448
|
+
"""
|
|
449
|
+
if n <= 0 or not probs:
|
|
450
|
+
return {}
|
|
451
|
+
items = list(probs.items())
|
|
452
|
+
raw = [(k, p * n) for k, p in items]
|
|
453
|
+
counts = {k: int(round(v)) for k, v in raw}
|
|
454
|
+
diff = n - sum(counts.values())
|
|
455
|
+
if diff != 0 and counts:
|
|
456
|
+
# Adjust the bucket with the largest fractional part
|
|
457
|
+
frac = sorted(raw, key=lambda kv: kv[1] - int(kv[1]), reverse=(diff > 0))
|
|
458
|
+
if frac:
|
|
459
|
+
counts[frac[0][0]] += diff
|
|
460
|
+
return counts
|
|
461
|
+
|
|
462
|
+
|
|
463
|
+
def compute_opinionsqa_ceiling(
|
|
464
|
+
data_dir: str | None = None, n_bootstrap: int = 1000
|
|
465
|
+
) -> dict | None:
|
|
466
|
+
"""Compute OpinionsQA ceiling from raw NONE_data.json counts (within-wave).
|
|
467
|
+
|
|
468
|
+
Returns a dict suitable for emission into leaderboard.json, or None if
|
|
469
|
+
raw data is not available on disk. Default n_bootstrap=1000 (the full
|
|
470
|
+
bootstrap budget); the vectorized JSD fast path in ``compute_ceiling``
|
|
471
|
+
keeps publish-time cost negligible at this B (sb-dkz).
|
|
472
|
+
"""
|
|
473
|
+
from pathlib import Path
|
|
474
|
+
|
|
475
|
+
from synthbench.datasets.opinionsqa import (
|
|
476
|
+
PEW_WAVES,
|
|
477
|
+
WAVE_YEAR_MAP,
|
|
478
|
+
_default_cache_dir,
|
|
479
|
+
)
|
|
480
|
+
|
|
481
|
+
data_path = Path(data_dir) if data_dir else _default_cache_dir()
|
|
482
|
+
human_resp = data_path / "raw" / "human_resp"
|
|
483
|
+
if not human_resp.is_dir():
|
|
484
|
+
return None
|
|
485
|
+
|
|
486
|
+
import json
|
|
487
|
+
|
|
488
|
+
wave_ceilings: list[CeilingResult] = []
|
|
489
|
+
wave_weights: list[float] = []
|
|
490
|
+
wave_details: list[dict] = []
|
|
491
|
+
|
|
492
|
+
for wave in PEW_WAVES:
|
|
493
|
+
wave_dir = human_resp / f"American_Trends_Panel_W{wave}"
|
|
494
|
+
none_path = wave_dir / "NONE_data.json"
|
|
495
|
+
if not none_path.exists():
|
|
496
|
+
continue
|
|
497
|
+
|
|
498
|
+
with open(none_path) as f:
|
|
499
|
+
data = json.load(f)
|
|
500
|
+
|
|
501
|
+
per_q_ceilings: list[CeilingResult] = []
|
|
502
|
+
for _qkey, entry in data.items():
|
|
503
|
+
if not isinstance(entry, dict):
|
|
504
|
+
continue
|
|
505
|
+
# Sum counts across sub_keys (political parties), per option
|
|
506
|
+
totals: dict[str, float] = {}
|
|
507
|
+
for sub_key, counts in entry.items():
|
|
508
|
+
if sub_key in ("MC_options", "question_text"):
|
|
509
|
+
continue
|
|
510
|
+
if not isinstance(counts, dict):
|
|
511
|
+
continue
|
|
512
|
+
for option, val in counts.items():
|
|
513
|
+
totals[option] = totals.get(option, 0.0) + float(val)
|
|
514
|
+
if not totals:
|
|
515
|
+
continue
|
|
516
|
+
int_counts = {k: int(round(v)) for k, v in totals.items() if v > 0}
|
|
517
|
+
if sum(int_counts.values()) < 10:
|
|
518
|
+
continue
|
|
519
|
+
try:
|
|
520
|
+
r = compute_ceiling_jsd(int_counts, n_bootstrap=n_bootstrap)
|
|
521
|
+
per_q_ceilings.append(r)
|
|
522
|
+
except ValueError:
|
|
523
|
+
continue
|
|
524
|
+
|
|
525
|
+
if per_q_ceilings:
|
|
526
|
+
agg = aggregate_ceilings(per_q_ceilings)
|
|
527
|
+
if agg is not None:
|
|
528
|
+
wave_ceilings.append(agg)
|
|
529
|
+
wave_weights.append(float(len(per_q_ceilings)))
|
|
530
|
+
wave_details.append(
|
|
531
|
+
{
|
|
532
|
+
"wave": f"ATP W{wave}",
|
|
533
|
+
"year": WAVE_YEAR_MAP.get(wave, 0),
|
|
534
|
+
"n_questions": len(per_q_ceilings),
|
|
535
|
+
"ceiling": agg.to_dict(),
|
|
536
|
+
}
|
|
537
|
+
)
|
|
538
|
+
|
|
539
|
+
if not wave_ceilings:
|
|
540
|
+
return None
|
|
541
|
+
|
|
542
|
+
overall = aggregate_ceilings(wave_ceilings, weights=wave_weights)
|
|
543
|
+
return {
|
|
544
|
+
"dataset": "opinionsqa",
|
|
545
|
+
"overall": overall.to_dict() if overall else None,
|
|
546
|
+
"per_wave": wave_details,
|
|
547
|
+
"protocol": "within_wave_split_half_multinomial_bootstrap",
|
|
548
|
+
"n_bootstrap": n_bootstrap,
|
|
549
|
+
}
|
|
550
|
+
|
|
551
|
+
|
|
552
|
+
def compute_opinionsqa_subgroup_ceilings(
|
|
553
|
+
data_dir: str | None = None, n_bootstrap: int = 1000
|
|
554
|
+
) -> dict | None:
|
|
555
|
+
"""Compute per-(wave, attribute, group) ceilings for OpinionsQA.
|
|
556
|
+
|
|
557
|
+
The aggregate ceiling from compute_opinionsqa_ceiling() is wave-level (n
|
|
558
|
+
~= 4000-5000 per wave) and overstates the achievable ceiling for P_sub,
|
|
559
|
+
which is measured at (wave × attribute × group) granularity where
|
|
560
|
+
subgroup sizes are 50-500. This function computes the ceiling at the
|
|
561
|
+
same granularity as the metric it bounds.
|
|
562
|
+
|
|
563
|
+
Aggregates per-question ceilings into one ceiling per (wave, attribute,
|
|
564
|
+
group), then emits the distribution (min/p25/median/p75/max) plus the
|
|
565
|
+
five worst subgroups by name.
|
|
566
|
+
|
|
567
|
+
Quality flags follow Cochran (1977): high n≥400, medium 200≤n<400,
|
|
568
|
+
low n<200. Small subgroups are retained but flagged so callers can
|
|
569
|
+
filter if needed.
|
|
570
|
+
|
|
571
|
+
Returns None if raw data is not on disk.
|
|
572
|
+
"""
|
|
573
|
+
from pathlib import Path
|
|
574
|
+
|
|
575
|
+
from synthbench.datasets.opinionsqa import (
|
|
576
|
+
PEW_WAVES,
|
|
577
|
+
WAVE_YEAR_MAP,
|
|
578
|
+
_default_cache_dir,
|
|
579
|
+
)
|
|
580
|
+
|
|
581
|
+
data_path = Path(data_dir) if data_dir else _default_cache_dir()
|
|
582
|
+
human_resp = data_path / "raw" / "human_resp"
|
|
583
|
+
if not human_resp.is_dir():
|
|
584
|
+
return None
|
|
585
|
+
|
|
586
|
+
import json
|
|
587
|
+
|
|
588
|
+
# Per-subgroup files shipped by Pew ATP. NONE_data.json is excluded
|
|
589
|
+
# because it is the wave-aggregate (not a subgroup).
|
|
590
|
+
ATTRIBUTE_FILES = [
|
|
591
|
+
"EDUCATION",
|
|
592
|
+
"POLPARTY",
|
|
593
|
+
"POLIDEOLOGY",
|
|
594
|
+
"RACE",
|
|
595
|
+
"SEX",
|
|
596
|
+
"INCOME",
|
|
597
|
+
"AGE",
|
|
598
|
+
"CREGION",
|
|
599
|
+
"POLPARTY_SEX",
|
|
600
|
+
"POLPARTY_RACE",
|
|
601
|
+
"RACE_SEX",
|
|
602
|
+
]
|
|
603
|
+
|
|
604
|
+
per_subgroup_rows: list[dict] = []
|
|
605
|
+
|
|
606
|
+
for wave in PEW_WAVES:
|
|
607
|
+
wave_dir = human_resp / f"American_Trends_Panel_W{wave}"
|
|
608
|
+
if not wave_dir.is_dir():
|
|
609
|
+
continue
|
|
610
|
+
year = WAVE_YEAR_MAP.get(wave, 0)
|
|
611
|
+
wave_label = f"ATP W{wave}"
|
|
612
|
+
|
|
613
|
+
for attr in ATTRIBUTE_FILES:
|
|
614
|
+
path = wave_dir / f"{attr}_data.json"
|
|
615
|
+
if not path.exists():
|
|
616
|
+
continue
|
|
617
|
+
|
|
618
|
+
with open(path) as f:
|
|
619
|
+
data = json.load(f)
|
|
620
|
+
|
|
621
|
+
# Bucket per-question ceilings by subgroup key. Cross-cut files
|
|
622
|
+
# (e.g. POLPARTY_SEX) ship one level of nesting — leaves are the
|
|
623
|
+
# option-count dicts we care about; branches become compound
|
|
624
|
+
# group labels like "Democrat × Female".
|
|
625
|
+
def _iter_subgroups(obj, prefix: str = ""):
|
|
626
|
+
for sub_key, sub_val in obj.items():
|
|
627
|
+
if sub_key in ("MC_options", "question_text"):
|
|
628
|
+
continue
|
|
629
|
+
if not isinstance(sub_val, dict):
|
|
630
|
+
continue
|
|
631
|
+
inner_vals = list(sub_val.values())
|
|
632
|
+
if inner_vals and all(isinstance(v, dict) for v in inner_vals):
|
|
633
|
+
# Nested cross-cut: recurse one level.
|
|
634
|
+
label = f"{prefix}{sub_key} × "
|
|
635
|
+
yield from _iter_subgroups(sub_val, prefix=label)
|
|
636
|
+
else:
|
|
637
|
+
yield f"{prefix}{sub_key}", sub_val
|
|
638
|
+
|
|
639
|
+
group_ceilings: dict[str, list[CeilingResult]] = {}
|
|
640
|
+
for _qkey, entry in data.items():
|
|
641
|
+
if not isinstance(entry, dict):
|
|
642
|
+
continue
|
|
643
|
+
for group_label, counts in _iter_subgroups(entry):
|
|
644
|
+
int_counts: dict[str, int] = {}
|
|
645
|
+
for k, v in counts.items():
|
|
646
|
+
try:
|
|
647
|
+
fv = float(v)
|
|
648
|
+
except (TypeError, ValueError):
|
|
649
|
+
continue
|
|
650
|
+
if fv > 0:
|
|
651
|
+
int_counts[k] = int(round(fv))
|
|
652
|
+
if sum(int_counts.values()) < 10:
|
|
653
|
+
continue
|
|
654
|
+
try:
|
|
655
|
+
r = compute_ceiling_jsd(int_counts, n_bootstrap=n_bootstrap)
|
|
656
|
+
except ValueError:
|
|
657
|
+
continue
|
|
658
|
+
group_ceilings.setdefault(group_label, []).append(r)
|
|
659
|
+
|
|
660
|
+
for group, results in sorted(group_ceilings.items()):
|
|
661
|
+
agg = aggregate_ceilings(results)
|
|
662
|
+
if agg is None:
|
|
663
|
+
continue
|
|
664
|
+
per_subgroup_rows.append(
|
|
665
|
+
{
|
|
666
|
+
"wave": wave_label,
|
|
667
|
+
"year": year,
|
|
668
|
+
"attribute": attr,
|
|
669
|
+
"group": group,
|
|
670
|
+
"n_questions": len(results),
|
|
671
|
+
"ceiling": agg.to_dict(),
|
|
672
|
+
}
|
|
673
|
+
)
|
|
674
|
+
|
|
675
|
+
if not per_subgroup_rows:
|
|
676
|
+
return None
|
|
677
|
+
|
|
678
|
+
values = np.array(
|
|
679
|
+
[row["ceiling"]["mean"] for row in per_subgroup_rows], dtype=np.float64
|
|
680
|
+
)
|
|
681
|
+
distribution = {
|
|
682
|
+
"min": round(float(values.min()), 6),
|
|
683
|
+
"p25": round(float(np.percentile(values, 25)), 6),
|
|
684
|
+
"median": round(float(np.percentile(values, 50)), 6),
|
|
685
|
+
"p75": round(float(np.percentile(values, 75)), 6),
|
|
686
|
+
"max": round(float(values.max()), 6),
|
|
687
|
+
}
|
|
688
|
+
|
|
689
|
+
worst_5 = sorted(per_subgroup_rows, key=lambda r: r["ceiling"]["mean"])[:5]
|
|
690
|
+
|
|
691
|
+
quality_breakdown: dict[str, int] = {"high": 0, "medium": 0, "low": 0}
|
|
692
|
+
for row in per_subgroup_rows:
|
|
693
|
+
flag = row["ceiling"]["quality_flag"]
|
|
694
|
+
quality_breakdown[flag] = quality_breakdown.get(flag, 0) + 1
|
|
695
|
+
|
|
696
|
+
return {
|
|
697
|
+
"dataset": "opinionsqa",
|
|
698
|
+
"granularity": "wave_attribute_group",
|
|
699
|
+
"distribution": distribution,
|
|
700
|
+
"subgroup_ceiling_for_psub": distribution["median"],
|
|
701
|
+
"worst_5_subgroups": worst_5,
|
|
702
|
+
"quality_breakdown": quality_breakdown,
|
|
703
|
+
"n_subgroups": len(per_subgroup_rows),
|
|
704
|
+
"per_subgroup": per_subgroup_rows,
|
|
705
|
+
"protocol": "per_subgroup_split_half_multinomial_bootstrap",
|
|
706
|
+
"n_bootstrap": n_bootstrap,
|
|
707
|
+
"note": (
|
|
708
|
+
"P_sub is a per-(wave, attribute, group) metric; the wave-aggregate "
|
|
709
|
+
"ceiling overstates achievable headroom at subgroup granularity. "
|
|
710
|
+
"Use the median subgroup ceiling (subgroup_ceiling_for_psub) as the "
|
|
711
|
+
"reference for P_sub, and the distribution to characterize spread."
|
|
712
|
+
),
|
|
713
|
+
}
|
|
714
|
+
|
|
715
|
+
|
|
716
|
+
def compute_subpop_ceiling(
|
|
717
|
+
data_dir: str | None = None,
|
|
718
|
+
n_per_subpop: int = 500,
|
|
719
|
+
n_bootstrap: int = 1000,
|
|
720
|
+
) -> dict | None:
|
|
721
|
+
"""Compute SubPOP ceiling per (attribute, group) subpopulation.
|
|
722
|
+
|
|
723
|
+
SubPOP ships probabilities without raw counts. We approximate counts by
|
|
724
|
+
rounding probs * n_per_subpop. Pew ATP subpops typically have
|
|
725
|
+
n ~= 300-800 respondents; the methodology writeup recommends flagging
|
|
726
|
+
ceilings as "medium" quality when n is inferred.
|
|
727
|
+
|
|
728
|
+
Returns a dict suitable for leaderboard.json, or None if raw data is
|
|
729
|
+
unavailable.
|
|
730
|
+
"""
|
|
731
|
+
from pathlib import Path
|
|
732
|
+
|
|
733
|
+
from synthbench.datasets.subpop import _default_cache_dir
|
|
734
|
+
|
|
735
|
+
data_path = Path(data_dir) if data_dir else _default_cache_dir()
|
|
736
|
+
raw_path = data_path / "raw_rows.json"
|
|
737
|
+
if not raw_path.exists():
|
|
738
|
+
return None
|
|
739
|
+
|
|
740
|
+
import json
|
|
741
|
+
|
|
742
|
+
with open(raw_path) as f:
|
|
743
|
+
rows = json.load(f)
|
|
744
|
+
|
|
745
|
+
# Group by (attribute, group)
|
|
746
|
+
subpop_ceilings: dict[tuple[str, str], list[CeilingResult]] = {}
|
|
747
|
+
for row in rows:
|
|
748
|
+
attr = row.get("attribute", "")
|
|
749
|
+
group = row.get("group", "")
|
|
750
|
+
options = row.get("options", [])
|
|
751
|
+
responses = row.get("responses", [])
|
|
752
|
+
if not attr or not group or len(options) != len(responses):
|
|
753
|
+
continue
|
|
754
|
+
probs = {opt: float(p) for opt, p in zip(options, responses)}
|
|
755
|
+
counts = _counts_from_probs(probs, n_per_subpop)
|
|
756
|
+
if sum(counts.values()) < 10:
|
|
757
|
+
continue
|
|
758
|
+
try:
|
|
759
|
+
r = compute_ceiling_jsd(counts, n_bootstrap=n_bootstrap)
|
|
760
|
+
subpop_ceilings.setdefault((attr, group), []).append(r)
|
|
761
|
+
except ValueError:
|
|
762
|
+
continue
|
|
763
|
+
|
|
764
|
+
if not subpop_ceilings:
|
|
765
|
+
return None
|
|
766
|
+
|
|
767
|
+
per_subpop: list[dict] = []
|
|
768
|
+
all_ceilings: list[CeilingResult] = []
|
|
769
|
+
all_weights: list[float] = []
|
|
770
|
+
for (attr, group), results in sorted(subpop_ceilings.items()):
|
|
771
|
+
agg = aggregate_ceilings(results)
|
|
772
|
+
if agg is None:
|
|
773
|
+
continue
|
|
774
|
+
# Downgrade flag: n is inferred, so cap at "medium"
|
|
775
|
+
if agg.quality_flag == "high":
|
|
776
|
+
agg = CeilingResult(
|
|
777
|
+
mean=agg.mean,
|
|
778
|
+
ci_low=agg.ci_low,
|
|
779
|
+
ci_high=agg.ci_high,
|
|
780
|
+
n_effective=agg.n_effective,
|
|
781
|
+
quality_flag="medium",
|
|
782
|
+
method=agg.method + "_inferred_n",
|
|
783
|
+
)
|
|
784
|
+
per_subpop.append(
|
|
785
|
+
{
|
|
786
|
+
"attribute": attr,
|
|
787
|
+
"group": group,
|
|
788
|
+
"n_questions": len(results),
|
|
789
|
+
"ceiling": agg.to_dict(),
|
|
790
|
+
}
|
|
791
|
+
)
|
|
792
|
+
all_ceilings.append(agg)
|
|
793
|
+
all_weights.append(float(len(results)))
|
|
794
|
+
|
|
795
|
+
overall = aggregate_ceilings(all_ceilings, weights=all_weights)
|
|
796
|
+
return {
|
|
797
|
+
"dataset": "subpop",
|
|
798
|
+
"overall": overall.to_dict() if overall else None,
|
|
799
|
+
"per_subpop": per_subpop,
|
|
800
|
+
"protocol": "per_subpop_split_half_multinomial_bootstrap",
|
|
801
|
+
"n_bootstrap": n_bootstrap,
|
|
802
|
+
"n_per_subpop_assumed": n_per_subpop,
|
|
803
|
+
"note": (
|
|
804
|
+
"SubPOP ships probabilities without raw counts; counts inferred "
|
|
805
|
+
f"at n={n_per_subpop} per subpop (typical Pew ATP subgroup size). "
|
|
806
|
+
"Quality flag capped at 'medium' due to inferred n."
|
|
807
|
+
),
|
|
808
|
+
}
|
|
809
|
+
|
|
810
|
+
|
|
811
|
+
def compute_globalopinionqa_ceiling(
|
|
812
|
+
data_dir: str | None = None,
|
|
813
|
+
n_per_country: int = 1000,
|
|
814
|
+
n_bootstrap: int = 1000,
|
|
815
|
+
) -> dict | None:
|
|
816
|
+
"""Compute GlobalOpinionQA ceiling per-country with regional aggregates.
|
|
817
|
+
|
|
818
|
+
Weighted by actual (country, question) coverage, not hypothetical coverage.
|
|
819
|
+
Typical Pew Global Attitudes n per country = 1000-1500.
|
|
820
|
+
"""
|
|
821
|
+
from pathlib import Path
|
|
822
|
+
|
|
823
|
+
from synthbench.datasets.globalopinionqa import _default_cache_dir
|
|
824
|
+
|
|
825
|
+
data_path = Path(data_dir) if data_dir else _default_cache_dir()
|
|
826
|
+
cache_path = data_path / "questions.json"
|
|
827
|
+
if not cache_path.exists():
|
|
828
|
+
return None
|
|
829
|
+
|
|
830
|
+
import json
|
|
831
|
+
|
|
832
|
+
with open(cache_path) as f:
|
|
833
|
+
payload = json.load(f)
|
|
834
|
+
questions = payload.get("questions", [])
|
|
835
|
+
if not questions:
|
|
836
|
+
return None
|
|
837
|
+
|
|
838
|
+
country_ceilings: dict[str, list[CeilingResult]] = {}
|
|
839
|
+
for q in questions:
|
|
840
|
+
options = q.get("options", [])
|
|
841
|
+
selections = q.get("selections", {})
|
|
842
|
+
if not options or not selections:
|
|
843
|
+
continue
|
|
844
|
+
for country, probs in selections.items():
|
|
845
|
+
if len(probs) != len(options):
|
|
846
|
+
continue
|
|
847
|
+
dist = {str(opt): float(p) for opt, p in zip(options, probs)}
|
|
848
|
+
counts = _counts_from_probs(dist, n_per_country)
|
|
849
|
+
if sum(counts.values()) < 10:
|
|
850
|
+
continue
|
|
851
|
+
try:
|
|
852
|
+
r = compute_ceiling_jsd(counts, n_bootstrap=n_bootstrap)
|
|
853
|
+
country_ceilings.setdefault(country, []).append(r)
|
|
854
|
+
except ValueError:
|
|
855
|
+
continue
|
|
856
|
+
|
|
857
|
+
if not country_ceilings:
|
|
858
|
+
return None
|
|
859
|
+
|
|
860
|
+
per_country: list[dict] = []
|
|
861
|
+
all_ceilings: list[CeilingResult] = []
|
|
862
|
+
all_weights: list[float] = []
|
|
863
|
+
for country, results in sorted(country_ceilings.items()):
|
|
864
|
+
agg = aggregate_ceilings(results)
|
|
865
|
+
if agg is None:
|
|
866
|
+
continue
|
|
867
|
+
# Flag as inferred n
|
|
868
|
+
if agg.quality_flag == "high":
|
|
869
|
+
agg = CeilingResult(
|
|
870
|
+
mean=agg.mean,
|
|
871
|
+
ci_low=agg.ci_low,
|
|
872
|
+
ci_high=agg.ci_high,
|
|
873
|
+
n_effective=agg.n_effective,
|
|
874
|
+
quality_flag="medium",
|
|
875
|
+
method=agg.method + "_inferred_n",
|
|
876
|
+
)
|
|
877
|
+
per_country.append(
|
|
878
|
+
{
|
|
879
|
+
"country": country,
|
|
880
|
+
"n_questions": len(results),
|
|
881
|
+
"ceiling": agg.to_dict(),
|
|
882
|
+
}
|
|
883
|
+
)
|
|
884
|
+
all_ceilings.append(agg)
|
|
885
|
+
# Weight aggregate by actual (country, question) coverage
|
|
886
|
+
all_weights.append(float(len(results)))
|
|
887
|
+
|
|
888
|
+
overall = aggregate_ceilings(all_ceilings, weights=all_weights)
|
|
889
|
+
return {
|
|
890
|
+
"dataset": "globalopinionqa",
|
|
891
|
+
"overall": overall.to_dict() if overall else None,
|
|
892
|
+
"per_country": per_country,
|
|
893
|
+
"protocol": "per_country_split_half_multinomial_bootstrap",
|
|
894
|
+
"n_bootstrap": n_bootstrap,
|
|
895
|
+
"n_per_country_assumed": n_per_country,
|
|
896
|
+
"note": (
|
|
897
|
+
"GlobalOpinionQA ships country probabilities without raw counts; "
|
|
898
|
+
f"counts inferred at n={n_per_country} per country (typical Pew "
|
|
899
|
+
"Global Attitudes survey size). Aggregate weighted by actual "
|
|
900
|
+
"(country, question) coverage."
|
|
901
|
+
),
|
|
902
|
+
}
|