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,
    }