File size: 5,358 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
"""Paired, replayable episodes with explicit failures and decision timings."""

import math
import statistics
import time
from collections.abc import Callable, Mapping, Sequence
from dataclasses import asdict
from typing import Any

from stackcraft.engine import new_game, step
from stackcraft.players import Player, observe, validate_decision
from stackcraft.replay import make_replay
from stackcraft.schema import RULES_VERSION


def _timing(seconds: list[float]) -> dict[str, float | int]:
    ordered = sorted(seconds)
    return {
        "count": len(seconds),
        "total_seconds": sum(seconds),
        "mean_seconds": statistics.mean(seconds) if seconds else 0.0,
        "median_seconds": statistics.median(seconds) if seconds else 0.0,
        "p95_seconds": ordered[math.ceil(0.95 * len(ordered)) - 1] if ordered else 0.0,
    }


def run_episode(player: Player, seed: int, max_pieces: int = 100) -> dict[str, Any]:
    """Stop at top-out, cap, or first error; never replace a failed player move."""
    if type(max_pieces) is not int or max_pieces < 1:
        raise ValueError("max_pieces must be a positive integer")
    state = new_game(seed)
    actions: list[str] = []
    decisions: list[dict[str, Any]] = []
    errors: list[dict[str, Any]] = []
    seconds: list[float] = []
    while not state.terminal and state.piece_index < max_pieces:
        observation = observe(state)
        event: dict[str, Any] = {"turn": state.piece_index, "observation": asdict(observation)}
        started = time.perf_counter()
        try:
            decision = player.choose(observation)
        except Exception as error:
            duration = time.perf_counter() - started
            failure = {"turn": state.piece_index, "kind": "player_error", "message": str(error)}
            errors.append(failure)
            event["error"] = failure
        else:
            duration = time.perf_counter() - started
            try:
                validate_decision(decision, observation)
            except ValueError as error:
                failure = {
                    "turn": state.piece_index,
                    "kind": "invalid_decision",
                    "message": str(error),
                }
                errors.append(failure)
                event["error"] = failure
            else:
                event["action_id"] = decision.action_id
                event["probabilities"] = decision.probabilities
                if (input_tokens := getattr(player, "last_input_tokens", None)) is not None:
                    event["input_tokens"] = input_tokens
                state = step(state, decision.action_id).state
                actions.append(decision.action_id)
        event["decision_seconds"] = duration
        seconds.append(duration)
        decisions.append(event)
        if errors:
            break
    cap_hit = not state.terminal and not errors and state.piece_index == max_pieces
    return {
        "schema_version": 1,
        "rules_version": RULES_VERSION,
        "player": {
            "name": player.name,
            "revision": player.revision,
            "runtime_config": getattr(player, "runtime_config", {}),
        },
        "seed": seed,
        "max_pieces": max_pieces,
        "outcome": {
            "score": state.score,
            "lines": state.lines,
            "pieces": state.piece_index,
            "terminal": state.terminal,
            "cap_hit": cap_hit,
            "status": "error" if errors else "top_out" if state.terminal else "cap",
        },
        "errors": errors,
        "decisions": decisions,
        "timing": _timing(seconds),
        "replay": make_replay(seed, actions),
    }


def run_tournament(
    player_factories: Mapping[str, Callable[[], Player]],
    seeds: Sequence[int],
    max_pieces: int = 100,
) -> dict[str, Any]:
    """Fresh player per paired seed. Timings are not deterministic game content."""
    seeds = tuple(seeds)
    if not player_factories or not seeds:
        raise ValueError("tournament requires players and seeds")
    if any(type(seed) is not int for seed in seeds) or len(set(seeds)) != len(seeds):
        raise ValueError("tournament seeds must be distinct integers")
    episodes = []
    summaries = {}
    for label, factory in player_factories.items():
        group = [run_episode(factory(), seed, max_pieces) for seed in seeds]
        for episode in group:
            episode["player_id"] = label
        episodes.extend(group)
        summaries[label] = {
            "episodes": len(group),
            "mean_lines": statistics.mean(e["outcome"]["lines"] for e in group),
            "median_lines": statistics.median(e["outcome"]["lines"] for e in group),
            "mean_score": statistics.mean(e["outcome"]["score"] for e in group),
            "mean_pieces": statistics.mean(e["outcome"]["pieces"] for e in group),
            "cap_hit_rate": statistics.mean(e["outcome"]["cap_hit"] for e in group),
            "error_rate": statistics.mean(bool(e["errors"]) for e in group),
            "timing": _timing([d["decision_seconds"] for e in group for d in e["decisions"]]),
        }
    return {
        "schema_version": 1,
        "rules_version": RULES_VERSION,
        "seeds": list(seeds),
        "max_pieces": max_pieces,
        "summary": summaries,
        "episodes": episodes,
    }