File size: 13,324 Bytes
4be6a52 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 | """Pure paired evaluation; never chooses checkpoints or runs game episodes."""
import json
import math
import random
import statistics
from collections import defaultdict
from collections.abc import Mapping, Sequence
from typing import Any
from stackcraft.data import canonical_json
from stackcraft.players import Decision, observe, validate_decision
from stackcraft.schema import RULES_VERSION, GameState
FINAL_TEST_SEEDS = tuple(range(30000, 30200))
METRICS = ("lines", "score", "pieces")
def _percentile(values: Sequence[float], quantile: float) -> float:
ordered = sorted(values)
position = (len(ordered) - 1) * quantile
lower = math.floor(position)
upper = math.ceil(position)
return ordered[lower] + (ordered[upper] - ordered[lower]) * (position - lower)
def paired_bootstrap(
differences: Sequence[float], *, samples: int = 10000, seed: int = 2026
) -> dict[str, Any]:
"""Percentile 95% interval, resampling complete paired episode differences."""
if not differences or any(
type(value) not in (int, float) or not math.isfinite(value) for value in differences
):
raise ValueError("paired differences must be nonempty finite numbers")
if type(samples) is not int or samples < 1 or type(seed) is not int:
raise ValueError("bootstrap samples must be positive and seed an integer")
rng = random.Random(seed)
count = len(differences)
means = [
sum(differences[rng.randrange(count)] for _ in range(count)) / count for _ in range(samples)
]
return {
"episodes": count,
"mean_difference": statistics.mean(differences),
"ci95_lower": _percentile(means, 0.025),
"ci95_upper": _percentile(means, 0.975),
"method": "paired episode percentile bootstrap",
"bootstrap_samples": samples,
"bootstrap_seed": seed,
}
def _numeric_summary(values: Sequence[float]) -> dict[str, Any]:
if not values:
return {"count": 0, "mean": None, "median": None, "p95": None}
return {
"count": len(values),
"mean": statistics.mean(values),
"median": statistics.median(values),
"p95": sorted(values)[math.ceil(0.95 * len(values)) - 1],
}
def _validate_episode(episode: dict[str, Any]) -> None:
if episode.get("schema_version") != 1 or type(episode.get("schema_version")) is not int:
raise ValueError("unsupported episode schema")
if episode.get("rules_version") != RULES_VERSION:
raise ValueError("unsupported episode rules")
if not isinstance(episode.get("player_id"), str) or not episode["player_id"]:
raise ValueError("each episode needs a nonempty player_id")
if type(episode.get("seed")) is not int:
raise ValueError("episode seed must be an integer")
cap = episode.get("max_pieces")
if type(cap) is not int or cap < 1:
raise ValueError("episode cap must be a positive integer")
player = episode.get("player")
if not isinstance(player, dict) or any(
not isinstance(player.get(key), str) or not player[key] for key in ("name", "revision")
):
raise ValueError("player name and revision are required")
if not isinstance(player.get("runtime_config", {}), dict):
raise ValueError("runtime configuration must be an object")
outcome = episode.get("outcome")
if not isinstance(outcome, dict) or any(
type(outcome.get(metric)) is not int or outcome[metric] < 0 for metric in METRICS
):
raise ValueError("outcome metrics must be nonnegative integers")
if outcome["pieces"] > cap:
raise ValueError("outcome exceeds piece cap")
if any(type(outcome.get(field)) is not bool for field in ("terminal", "cap_hit")):
raise ValueError("terminal and cap_hit must be booleans")
errors, decisions = episode.get("errors"), episode.get("decisions")
if not isinstance(errors, list) or not isinstance(decisions, list):
raise ValueError("episode errors and decisions must be lists")
if len(errors) > 1 or any(
not isinstance(error, dict) or error.get("kind") not in ("player_error", "invalid_decision")
for error in errors
):
raise ValueError("invalid episode error record")
expected_status = "error" if errors else "top_out" if outcome["terminal"] else "cap"
expected_cap = not errors and not outcome["terminal"] and outcome["pieces"] == cap
if outcome.get("status") != expected_status or outcome["cap_hit"] != expected_cap:
raise ValueError("outcome status is inconsistent with errors, terminal state, or cap")
if not errors and not outcome["terminal"] and outcome["pieces"] != cap:
raise ValueError("successful nonterminal episode ended before its cap")
if errors and outcome["terminal"]:
raise ValueError("an errored episode cannot also finish in top-out")
if len(decisions) != outcome["pieces"] + bool(errors):
raise ValueError("decision count does not match placements and errors")
for event in decisions:
duration = event.get("decision_seconds") if isinstance(event, dict) else None
if not isinstance(duration, (int, float)) or isinstance(duration, bool):
raise ValueError("decision duration must be a finite nonnegative number")
if not math.isfinite(duration) or duration < 0:
raise ValueError("decision duration must be a finite nonnegative number")
tokens = event.get("input_tokens")
if tokens is not None and (type(tokens) is not int or tokens < 1):
raise ValueError("input_tokens must be a positive integer when provided")
def paired_report(
episodes: Sequence[dict[str, Any]],
*,
trained_id: str,
base_id: str,
bootstrap_samples: int = 10000,
bootstrap_seed: int = 2026,
max_error_rate: float = 0.0,
) -> dict[str, Any]:
"""Failure-adjusted primary outcomes keep failed episodes at zero, never drop them.
Observed partial outcomes are also reported. A positive interval is statistical
evidence for this comparison, not automatic checkpoint selection or publication.
"""
if not episodes or trained_id == base_id:
raise ValueError("report requires episodes and distinct trained/base players")
if type(max_error_rate) not in (int, float) or not 0 <= max_error_rate <= 1:
raise ValueError("max_error_rate must be between zero and one")
grouped: dict[str, dict[int, dict[str, Any]]] = defaultdict(dict)
identities: dict[str, str] = {}
caps = set()
for episode in episodes:
_validate_episode(episode)
player_id, seed = episode["player_id"], episode["seed"]
if seed in grouped[player_id]:
raise ValueError("duplicate player/seed episode")
grouped[player_id][seed] = episode
caps.add(episode["max_pieces"])
identity = canonical_json(episode["player"])
if identities.setdefault(player_id, identity) != identity:
raise ValueError("player revision or runtime changed within evaluation")
if trained_id not in grouped or base_id not in grouped:
raise ValueError("both requested players must be present")
seeds = sorted(grouped[base_id])
if len(caps) != 1 or any(sorted(group) != seeds for group in grouped.values()):
raise ValueError("all players must have identical seed sets and episode caps")
summary: dict[str, Any] = {}
for player_id, group in grouped.items():
ordered = [group[seed] for seed in seeds]
failed = sum(bool(episode["errors"]) for episode in ordered)
decisions = [event for episode in ordered for event in episode["decisions"]]
primary = {
metric: [0 if e["errors"] else e["outcome"][metric] for e in ordered]
for metric in METRICS
}
summary[player_id] = {
"identity": json.loads(identities[player_id]),
"episodes": len(ordered),
"failed_episodes": failed,
"error_rate": failed / len(ordered),
"invalid_decisions": sum(
error["kind"] == "invalid_decision" for e in ordered for error in e["errors"]
),
"cap_hit_rate": statistics.mean(e["outcome"]["cap_hit"] for e in ordered),
"failure_adjusted": {
metric: _numeric_summary(values) for metric, values in primary.items()
},
"observed_before_error": {
metric: _numeric_summary([e["outcome"][metric] for e in ordered])
for metric in METRICS
},
"latency_seconds": _numeric_summary([d["decision_seconds"] for d in decisions]),
"input_tokens": _numeric_summary(
[d["input_tokens"] for d in decisions if d.get("input_tokens") is not None]
),
"failed_seeds": [e["seed"] for e in ordered if e["errors"]],
}
comparison = {}
for metric in METRICS:
differences = []
for seed in seeds:
trained, base = grouped[trained_id][seed], grouped[base_id][seed]
difference = (0 if trained["errors"] else trained["outcome"][metric]) - (
0 if base["errors"] else base["outcome"][metric]
)
differences.append(difference)
comparison[metric] = paired_bootstrap(
differences, samples=bootstrap_samples, seed=bootstrap_seed
)
errors_acceptable = all(
summary[player]["error_rate"] <= max_error_rate for player in (trained_id, base_id)
)
return {
"schema_version": 1,
"rules_version": RULES_VERSION,
"seeds": seeds,
"matches_reserved_test_pool": tuple(seeds) == FINAL_TEST_SEEDS,
"max_pieces": next(iter(caps)),
"trained_id": trained_id,
"base_id": base_id,
"primary_failure_policy": "zero lines/score/pieces on errors; retain all pairs",
"players": summary,
"paired_trained_minus_base": comparison,
"max_error_rate": max_error_rate,
"errors_acceptable": errors_acceptable,
"positive_mean_lines_signal": (
len(seeds) >= 2 and comparison["lines"]["ci95_lower"] > 0 and errors_acceptable
),
"selection_performed": False,
}
def summarize_positions(
records: Sequence[dict[str, Any]], predictions: Mapping[str, Decision | None]
) -> dict[str, Any]:
"""Teacher agreement and proper probability scores, separately from game outcomes.
Missing/invalid predictions count against agreement. NLL/Brier require complete
probabilities; if any prediction fails or omits them, aggregate scores are null.
Finite-subset diagnostic scores are labeled with their smaller denominator.
"""
ids = [row["id"] for row in records]
if not ids or len(set(ids)) != len(ids) or set(predictions) - set(ids):
raise ValueError("position IDs must be nonempty, unique, and match prediction keys")
correct = 0
failures = []
missing_probabilities = []
nll = []
brier = []
for row in records:
raw = row["observation"]
state = GameState(
tuple(tuple(r) for r in raw["board"]), 0, 0, raw["current"], raw["next_piece"]
)
observation = observe(state)
target = row["action_id"]
if target not in {action.id for action in observation.legal_actions}:
raise ValueError("teacher label is illegal")
decision = predictions.get(row["id"])
try:
if decision is None:
raise ValueError("missing prediction")
validate_decision(decision, observation)
except ValueError:
failures.append(row["id"])
continue
correct += decision.action_id == target
probabilities = decision.probabilities
if probabilities is None:
missing_probabilities.append(row["id"])
continue
# A stated zero target probability means infinite NLL; preserve that fact
# explicitly without putting nonstandard Infinity values into JSON.
nll.append(-math.log(probabilities[target]) if probabilities[target] > 0 else None)
brier.append(sum((p - int(key == target)) ** 2 for key, p in probabilities.items()))
complete = len(brier) == len(records)
finite_nll = [value for value in nll if value is not None]
zero_count = len(nll) - len(finite_nll)
return {
"positions": len(records),
"correct": correct,
"teacher_agreement": correct / len(records),
"failed_predictions": failures,
"error_rate": len(failures) / len(records),
"missing_probabilities": missing_probabilities,
"probability_positions": len(brier),
"complete_probability_coverage": complete,
"mean_nll": statistics.mean(finite_nll) if complete and not zero_count else None,
"nll_is_infinite": bool(zero_count),
"zero_target_probability_count": zero_count,
"mean_brier": statistics.mean(brier) if complete else None,
"finite_subset_mean_nll": statistics.mean(finite_nll) if finite_nll else None,
"finite_subset_nll_positions": len(finite_nll),
"available_subset_mean_brier": statistics.mean(brier) if brier else None,
"selection_performed": False,
}
|