nashbench 0.1.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.
- nashbench/__init__.py +29 -0
- nashbench/core.py +216 -0
- nashbench/exploitability.py +90 -0
- nashbench/games/__init__.py +38 -0
- nashbench/games/_goofspiel_tree.py +240 -0
- nashbench/games/_phantom.py +114 -0
- nashbench/games/_phantom_tree.py +399 -0
- nashbench/games/dark_hex.py +49 -0
- nashbench/games/goofspiel.py +117 -0
- nashbench/games/kuhn_poker.py +89 -0
- nashbench/games/leduc_poker.py +120 -0
- nashbench/games/phantom_ttt.py +30 -0
- nashbench/policy.py +38 -0
- nashbench/py.typed +0 -0
- nashbench/sequence_form.py +214 -0
- nashbench-0.1.0.dist-info/METADATA +102 -0
- nashbench-0.1.0.dist-info/RECORD +19 -0
- nashbench-0.1.0.dist-info/WHEEL +4 -0
- nashbench-0.1.0.dist-info/licenses/LICENSE +202 -0
nashbench/__init__.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""Benchmark games and evaluation for Nash equilibrium solvers."""
|
|
2
|
+
|
|
3
|
+
import importlib.metadata
|
|
4
|
+
|
|
5
|
+
from nashbench.core import BOTH
|
|
6
|
+
from nashbench.core import Game
|
|
7
|
+
from nashbench.core import State
|
|
8
|
+
from nashbench.core import TimeStep
|
|
9
|
+
from nashbench.core import auto_reset
|
|
10
|
+
from nashbench.exploitability import exploitability
|
|
11
|
+
from nashbench.games import REGISTRY
|
|
12
|
+
from nashbench.games import make
|
|
13
|
+
from nashbench.policy import Policy
|
|
14
|
+
from nashbench.policy import uniform_random
|
|
15
|
+
|
|
16
|
+
__version__ = importlib.metadata.version("nashbench")
|
|
17
|
+
|
|
18
|
+
__all__ = [
|
|
19
|
+
"BOTH",
|
|
20
|
+
"REGISTRY",
|
|
21
|
+
"Game",
|
|
22
|
+
"Policy",
|
|
23
|
+
"State",
|
|
24
|
+
"TimeStep",
|
|
25
|
+
"auto_reset",
|
|
26
|
+
"exploitability",
|
|
27
|
+
"make",
|
|
28
|
+
"uniform_random",
|
|
29
|
+
]
|
nashbench/core.py
ADDED
|
@@ -0,0 +1,216 @@
|
|
|
1
|
+
"""The game API: `Game`, its `State`, and the `TimeStep` players observe."""
|
|
2
|
+
|
|
3
|
+
import abc
|
|
4
|
+
from collections.abc import Callable
|
|
5
|
+
import dataclasses
|
|
6
|
+
import functools
|
|
7
|
+
from typing import dataclass_transform
|
|
8
|
+
|
|
9
|
+
import jax
|
|
10
|
+
import jax.numpy as jnp
|
|
11
|
+
|
|
12
|
+
from nashbench import sequence_form as sequence_form_lib
|
|
13
|
+
|
|
14
|
+
#: `TimeStep.current_player` when both players act at once. Equals OpenSpiel's
|
|
15
|
+
#: kSimultaneousPlayerId.
|
|
16
|
+
BOTH = -2
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass_transform(frozen_default=True)
|
|
20
|
+
def pytree_dataclass[T](cls: type[T]) -> type[T]:
|
|
21
|
+
"""Makes `cls` a frozen dataclass that JAX transformations accept."""
|
|
22
|
+
return jax.tree_util.register_dataclass(
|
|
23
|
+
dataclasses.dataclass(frozen=True)(cls)
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@pytree_dataclass
|
|
28
|
+
class TimeStep:
|
|
29
|
+
"""What the players see after `Game.reset` or `Game.step`.
|
|
30
|
+
|
|
31
|
+
Arrays with a leading axis of size 2 are indexed by player.
|
|
32
|
+
`observation` and `legal_action_mask` are defined only for the player to
|
|
33
|
+
act and only while `done` is false; other entries are unspecified.
|
|
34
|
+
`reward` is always defined for both players.
|
|
35
|
+
|
|
36
|
+
Attributes:
|
|
37
|
+
observation: `[2, *observation_shape]` float32 information states.
|
|
38
|
+
legal_action_mask: `[2, num_actions]` bool.
|
|
39
|
+
reward: `[2]` float32 reward of the last step.
|
|
40
|
+
done: `[]` bool, true once the episode has ended.
|
|
41
|
+
current_player: `[]` int32 player to act: 0, 1, or `BOTH`.
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
observation: jax.Array
|
|
45
|
+
legal_action_mask: jax.Array
|
|
46
|
+
reward: jax.Array
|
|
47
|
+
done: jax.Array
|
|
48
|
+
current_player: jax.Array
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@pytree_dataclass
|
|
52
|
+
class GameState:
|
|
53
|
+
"""Base class for the state of a game's rules, indexed by seat.
|
|
54
|
+
|
|
55
|
+
Attributes:
|
|
56
|
+
current_seat: `[]` int32 seat to act: 0, 1, or `BOTH`.
|
|
57
|
+
done: `[]` bool, true once the game has ended.
|
|
58
|
+
"""
|
|
59
|
+
|
|
60
|
+
current_seat: jax.Array
|
|
61
|
+
done: jax.Array
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@pytree_dataclass
|
|
65
|
+
class State[S: GameState]:
|
|
66
|
+
"""The full state of an episode, including what players cannot see.
|
|
67
|
+
|
|
68
|
+
Attributes:
|
|
69
|
+
key: PRNG key of the next episode, used by `auto_reset`.
|
|
70
|
+
seat: `[2]` int32 permutation; `seat[p]` is the seat of player `p`.
|
|
71
|
+
game_state: The state of the game's rules.
|
|
72
|
+
"""
|
|
73
|
+
|
|
74
|
+
key: jax.Array
|
|
75
|
+
seat: jax.Array
|
|
76
|
+
game_state: S
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class Game[S: GameState](abc.ABC):
|
|
80
|
+
"""A two-player zero-sum game.
|
|
81
|
+
|
|
82
|
+
Players interact with a game through `reset` and `step`. At every reset,
|
|
83
|
+
a fair coin assigns the two players to the game's two seats, where seat
|
|
84
|
+
`i` is OpenSpiel's player `i`.
|
|
85
|
+
|
|
86
|
+
To add a game, subclass `Game`, set `num_actions` and `observation_shape`,
|
|
87
|
+
and implement the rules: `initial_states`, `apply_action`, `observe`,
|
|
88
|
+
`legal_action_mask`, and `returns`. The rules deal with seats only;
|
|
89
|
+
`reset` and `step` translate between seats and players.
|
|
90
|
+
|
|
91
|
+
Attributes:
|
|
92
|
+
num_actions: Number of actions.
|
|
93
|
+
observation_shape: Shape of one player's observation.
|
|
94
|
+
"""
|
|
95
|
+
|
|
96
|
+
num_actions: int
|
|
97
|
+
observation_shape: tuple[int, ...]
|
|
98
|
+
|
|
99
|
+
def reset(self, key: jax.Array) -> tuple[State[S], TimeStep]:
|
|
100
|
+
"""Starts an episode, assigning players to seats at random."""
|
|
101
|
+
key, seat_key, chance_key = jax.random.split(key, 3)
|
|
102
|
+
states, probs = self.initial_states()
|
|
103
|
+
i = jax.random.choice(chance_key, probs.size, p=probs)
|
|
104
|
+
state = State(
|
|
105
|
+
key=key,
|
|
106
|
+
seat=jax.random.permutation(seat_key, 2),
|
|
107
|
+
game_state=jax.tree.map(lambda x: x[i], states),
|
|
108
|
+
)
|
|
109
|
+
return state, self._timestep(state, reward=jnp.zeros(2))
|
|
110
|
+
|
|
111
|
+
def step(
|
|
112
|
+
self, state: State[S], action: jax.Array
|
|
113
|
+
) -> tuple[State[S], TimeStep]:
|
|
114
|
+
"""Advances the episode by one step.
|
|
115
|
+
|
|
116
|
+
Args:
|
|
117
|
+
state: The current state.
|
|
118
|
+
action: `[2]` int32 action of each player. Actions of players who
|
|
119
|
+
are not acting are ignored.
|
|
120
|
+
|
|
121
|
+
Returns:
|
|
122
|
+
The next state and timestep. Once the episode is done, `step`
|
|
123
|
+
leaves the state unchanged and returns zero rewards.
|
|
124
|
+
"""
|
|
125
|
+
# A permutation of two elements is its own inverse, so indexing with
|
|
126
|
+
# `seat` maps player-indexed arrays to seat-indexed ones and back.
|
|
127
|
+
old = state.game_state
|
|
128
|
+
new = self.apply_action(old, action[state.seat])
|
|
129
|
+
new = jax.tree.map(lambda o, n: jnp.where(old.done, o, n), old, new)
|
|
130
|
+
reward = (self.returns(new) - self.returns(old))[state.seat]
|
|
131
|
+
state = dataclasses.replace(state, game_state=new)
|
|
132
|
+
return state, self._timestep(state, reward)
|
|
133
|
+
|
|
134
|
+
def _timestep(self, state: State[S], reward: jax.Array) -> TimeStep:
|
|
135
|
+
game_state = state.game_state
|
|
136
|
+
observe = jax.vmap(self.observe, in_axes=(None, 0))
|
|
137
|
+
legal_action_mask = jax.vmap(self.legal_action_mask, in_axes=(None, 0))
|
|
138
|
+
current_seat = game_state.current_seat
|
|
139
|
+
return TimeStep(
|
|
140
|
+
observation=observe(game_state, state.seat),
|
|
141
|
+
legal_action_mask=legal_action_mask(game_state, state.seat),
|
|
142
|
+
reward=reward.astype(jnp.float32),
|
|
143
|
+
done=game_state.done,
|
|
144
|
+
current_player=jnp.where(
|
|
145
|
+
current_seat == BOTH, BOTH, state.seat[current_seat]
|
|
146
|
+
),
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
@abc.abstractmethod
|
|
150
|
+
def initial_states(self) -> tuple[S, jax.Array]:
|
|
151
|
+
"""Returns every possible initial state, stacked, and its probability.
|
|
152
|
+
|
|
153
|
+
All chance events, such as card deals, happen at the start, so initial
|
|
154
|
+
states differ only in their outcomes.
|
|
155
|
+
"""
|
|
156
|
+
|
|
157
|
+
@abc.abstractmethod
|
|
158
|
+
def apply_action(self, state: S, action: jax.Array) -> S:
|
|
159
|
+
"""Returns the state after the seats take `action`.
|
|
160
|
+
|
|
161
|
+
Args:
|
|
162
|
+
state: The current state. If it is done, the result is discarded.
|
|
163
|
+
action: `[2]` int32 action of each seat. Turn-based games read
|
|
164
|
+
`action[state.current_seat]`.
|
|
165
|
+
"""
|
|
166
|
+
|
|
167
|
+
@abc.abstractmethod
|
|
168
|
+
def observe(self, state: S, seat: jax.Array) -> jax.Array:
|
|
169
|
+
"""Returns the information state of `seat` as a float32 tensor.
|
|
170
|
+
|
|
171
|
+
The tensor is a one-hot encoding of `seat` followed by OpenSpiel's
|
|
172
|
+
information-state tensor for that seat. If OpenSpiel's tensor already
|
|
173
|
+
starts with the seat, it is not repeated.
|
|
174
|
+
"""
|
|
175
|
+
|
|
176
|
+
@abc.abstractmethod
|
|
177
|
+
def legal_action_mask(self, state: S, seat: jax.Array) -> jax.Array:
|
|
178
|
+
"""Returns the `[num_actions]` bool mask of legal actions of `seat`."""
|
|
179
|
+
|
|
180
|
+
@abc.abstractmethod
|
|
181
|
+
def returns(self, state: S) -> jax.Array:
|
|
182
|
+
"""Returns the `[2]` total reward of each seat so far."""
|
|
183
|
+
|
|
184
|
+
@functools.cached_property
|
|
185
|
+
def sequence_form(self) -> sequence_form_lib.SequenceForm:
|
|
186
|
+
"""The sequence form of the game, for exact evaluation.
|
|
187
|
+
|
|
188
|
+
Built on first access and cached with the game. By default, enumerates
|
|
189
|
+
the game tree; games whose tree doesn't fit in memory override this.
|
|
190
|
+
"""
|
|
191
|
+
return sequence_form_lib.enumerate_tree(self)
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def auto_reset(
|
|
195
|
+
game: Game,
|
|
196
|
+
) -> Callable[[State, jax.Array], tuple[State, TimeStep]]:
|
|
197
|
+
"""Returns `game.step`, modified to start a new episode when one ends.
|
|
198
|
+
|
|
199
|
+
On the step that ends an episode, the timestep keeps that step's `reward`
|
|
200
|
+
and `done`. Its other fields, and the returned state, belong to the first
|
|
201
|
+
step of the next episode.
|
|
202
|
+
"""
|
|
203
|
+
|
|
204
|
+
def step(state: State, action: jax.Array) -> tuple[State, TimeStep]:
|
|
205
|
+
state, timestep = game.step(state, action)
|
|
206
|
+
next_state, next_timestep = game.reset(state.key)
|
|
207
|
+
next_timestep = dataclasses.replace(
|
|
208
|
+
next_timestep, reward=timestep.reward, done=timestep.done
|
|
209
|
+
)
|
|
210
|
+
return jax.tree.map(
|
|
211
|
+
lambda n, o: jnp.where(timestep.done, n, o),
|
|
212
|
+
(next_state, next_timestep),
|
|
213
|
+
(state, timestep),
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
return step
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
"""Exploitability: how far a policy is from a Nash equilibrium."""
|
|
2
|
+
|
|
3
|
+
import functools
|
|
4
|
+
|
|
5
|
+
import jax
|
|
6
|
+
import jax.numpy as jnp
|
|
7
|
+
|
|
8
|
+
from nashbench import core
|
|
9
|
+
from nashbench import policy as policy_lib
|
|
10
|
+
from nashbench import sequence_form
|
|
11
|
+
|
|
12
|
+
# Policies see information states in chunks of this size.
|
|
13
|
+
_CHUNK_SIZE = 2**16
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def exploitability(game: core.Game, policy: policy_lib.Policy) -> float:
|
|
17
|
+
"""Returns the exploitability of `policy` playing both seats of `game`.
|
|
18
|
+
|
|
19
|
+
Exploitability is the value a best response gains against `policy`,
|
|
20
|
+
averaged over the two seats:
|
|
21
|
+
|
|
22
|
+
(max_π u_0(π, policy) + max_π u_1(policy, π)) / 2,
|
|
23
|
+
|
|
24
|
+
where `u_i(σ_0, σ_1)` is the expected return of seat `i` when seat 0
|
|
25
|
+
plays `σ_0` and seat 1 plays `σ_1`. It is 0 exactly when `policy` is a
|
|
26
|
+
Nash equilibrium, and equals OpenSpiel's `exploitability.exploitability`
|
|
27
|
+
(NashConv / 2).
|
|
28
|
+
|
|
29
|
+
The first call for a game builds its sequence form, which takes up to a
|
|
30
|
+
minute for the largest games; later calls with the same game reuse it.
|
|
31
|
+
|
|
32
|
+
Args:
|
|
33
|
+
game: The game.
|
|
34
|
+
policy: The policy, playing both seats.
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
The exploitability, in units of the game's rewards.
|
|
38
|
+
|
|
39
|
+
Raises:
|
|
40
|
+
ValueError: If the policy doesn't return probabilities over legal
|
|
41
|
+
actions.
|
|
42
|
+
"""
|
|
43
|
+
form = game.sequence_form
|
|
44
|
+
probs = _action_probabilities(form, policy)
|
|
45
|
+
plan = _realization_plan(probs, form.parent, form.levels)
|
|
46
|
+
values = _best_response_values(
|
|
47
|
+
form.gradient(plan), form.parent, form.legal_action_mask, form.levels
|
|
48
|
+
)
|
|
49
|
+
return float(values.sum() / 2)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _action_probabilities(
|
|
53
|
+
form: sequence_form.SequenceForm, policy: policy_lib.Policy
|
|
54
|
+
) -> jax.Array:
|
|
55
|
+
"""Returns `[num_infosets, num_actions]` probabilities of `policy`."""
|
|
56
|
+
num_infosets = form.parent.size
|
|
57
|
+
chunks = []
|
|
58
|
+
for start in range(0, num_infosets, _CHUNK_SIZE):
|
|
59
|
+
ids = jnp.arange(start, min(start + _CHUNK_SIZE, num_infosets))
|
|
60
|
+
mask = form.legal_action_mask[ids]
|
|
61
|
+
chunks.append(jax.vmap(policy)(form.observe(ids), mask))
|
|
62
|
+
probs = jnp.concatenate(chunks)
|
|
63
|
+
illegal = jnp.where(form.legal_action_mask, 0.0, probs).max()
|
|
64
|
+
if illegal > 1e-6 or jnp.abs(probs.sum(1) - 1).max() > 1e-4:
|
|
65
|
+
raise ValueError(
|
|
66
|
+
"The policy must return probabilities of legal actions."
|
|
67
|
+
)
|
|
68
|
+
return probs
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
@functools.partial(jax.jit, static_argnames="levels")
|
|
72
|
+
def _realization_plan(probs, parent, levels):
|
|
73
|
+
"""Returns the realization plan of both seats, top-down by level."""
|
|
74
|
+
plan = jnp.ones((probs.shape[0] + 1, probs.shape[1]))
|
|
75
|
+
for start, end in levels:
|
|
76
|
+
reach = plan.ravel()[parent[start:end], None] * probs[start:end]
|
|
77
|
+
plan = plan.at[1 + start : 1 + end].set(reach)
|
|
78
|
+
return plan.ravel()
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
@functools.partial(jax.jit, static_argnames="levels")
|
|
82
|
+
def _best_response_values(gradient, parent, legal_action_mask, levels):
|
|
83
|
+
"""Returns the `[2]` value of each seat's best response, bottom-up."""
|
|
84
|
+
values = gradient.reshape(-1, legal_action_mask.shape[1])
|
|
85
|
+
for start, end in reversed(levels):
|
|
86
|
+
q = values[1 + start : 1 + end]
|
|
87
|
+
best = jnp.where(legal_action_mask[start:end], q, -jnp.inf).max(1)
|
|
88
|
+
values = values.ravel().at[parent[start:end]].add(best)
|
|
89
|
+
values = values.reshape(-1, legal_action_mask.shape[1])
|
|
90
|
+
return values[0, :2]
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
"""The benchmark games, by name."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Callable
|
|
4
|
+
import functools
|
|
5
|
+
|
|
6
|
+
from nashbench import core
|
|
7
|
+
from nashbench.games import dark_hex
|
|
8
|
+
from nashbench.games import goofspiel
|
|
9
|
+
from nashbench.games import kuhn_poker
|
|
10
|
+
from nashbench.games import leduc_poker
|
|
11
|
+
from nashbench.games import phantom_ttt
|
|
12
|
+
|
|
13
|
+
#: Constructor of each game, by name.
|
|
14
|
+
REGISTRY: dict[str, Callable[..., core.Game]] = {
|
|
15
|
+
"kuhn_poker": kuhn_poker.KuhnPoker,
|
|
16
|
+
"leduc_poker": leduc_poker.LeducPoker,
|
|
17
|
+
"phantom_ttt": functools.partial(phantom_ttt.PhantomTTT, abrupt=False),
|
|
18
|
+
"phantom_ttt_abrupt": functools.partial(
|
|
19
|
+
phantom_ttt.PhantomTTT, abrupt=True
|
|
20
|
+
),
|
|
21
|
+
"dark_hex3": functools.partial(dark_hex.DarkHex3, abrupt=False),
|
|
22
|
+
"dark_hex3_abrupt": functools.partial(dark_hex.DarkHex3, abrupt=True),
|
|
23
|
+
"goofspiel": goofspiel.Goofspiel,
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def make(name: str, **kwargs) -> core.Game:
|
|
28
|
+
"""Returns the game named `name`, a key of `REGISTRY`.
|
|
29
|
+
|
|
30
|
+
Args:
|
|
31
|
+
name: Name of the game.
|
|
32
|
+
**kwargs: Game parameters, such as `num_cards` of `"goofspiel"`.
|
|
33
|
+
"""
|
|
34
|
+
if name not in REGISTRY:
|
|
35
|
+
raise ValueError(
|
|
36
|
+
f"Unknown game {name!r}; choose from {list(REGISTRY)}."
|
|
37
|
+
)
|
|
38
|
+
return REGISTRY[name](**kwargs)
|
|
@@ -0,0 +1,240 @@
|
|
|
1
|
+
"""The sequence form of Goofspiel, from the regular structure of its tree.
|
|
2
|
+
|
|
3
|
+
At turn `t`, a seat's information set is the first `t + 1` point cards, its
|
|
4
|
+
first `t` bids, and who won each of the first `t` turns. Point cards are
|
|
5
|
+
independent of bids, so the information sets of turn `t` are numbered
|
|
6
|
+
|
|
7
|
+
offset[t] + rank(point cards) * num_pairs[t] + pair,
|
|
8
|
+
|
|
9
|
+
where `pair` numbers the (bids, winners) that some opponent bids make
|
|
10
|
+
possible. A terminal history is a triple of permutations: the point cards and
|
|
11
|
+
each seat's bids.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
import dataclasses
|
|
15
|
+
import functools
|
|
16
|
+
import itertools
|
|
17
|
+
import math
|
|
18
|
+
|
|
19
|
+
import jax
|
|
20
|
+
import jax.numpy as jnp
|
|
21
|
+
import numpy as np
|
|
22
|
+
|
|
23
|
+
from nashbench import sequence_form as sequence_form_lib
|
|
24
|
+
|
|
25
|
+
# Information sets per chunk when building tables.
|
|
26
|
+
_CHUNK = 2**22
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def sequence_form(game) -> sequence_form_lib.SequenceForm:
|
|
30
|
+
"""Returns the sequence form of a Goofspiel game."""
|
|
31
|
+
k = game.num_cards
|
|
32
|
+
tables = _Tables(k)
|
|
33
|
+
ids = jnp.arange(2 * tables.num_infosets)
|
|
34
|
+
parent, legal = (
|
|
35
|
+
jnp.concatenate(x)
|
|
36
|
+
for x in zip(
|
|
37
|
+
*(
|
|
38
|
+
_parent_and_legal(tables, ids[i : i + _CHUNK])
|
|
39
|
+
for i in range(0, ids.size, _CHUNK)
|
|
40
|
+
),
|
|
41
|
+
strict=True,
|
|
42
|
+
)
|
|
43
|
+
)
|
|
44
|
+
levels = tuple(
|
|
45
|
+
(
|
|
46
|
+
s * tables.num_infosets + int(tables.offset[t]),
|
|
47
|
+
s * tables.num_infosets + int(tables.offset[t + 1]),
|
|
48
|
+
)
|
|
49
|
+
for t in range(k - 1)
|
|
50
|
+
for s in range(2)
|
|
51
|
+
)
|
|
52
|
+
observe = jax.jit(functools.partial(_observe, game, tables))
|
|
53
|
+
return sequence_form_lib.SequenceForm(
|
|
54
|
+
levels=levels,
|
|
55
|
+
parent=parent,
|
|
56
|
+
legal_action_mask=legal,
|
|
57
|
+
observe=observe,
|
|
58
|
+
gradient=jax.jit(functools.partial(_gradient, tables)),
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
class _Tables:
|
|
63
|
+
"""Small tables that number Goofspiel's information sets and leaves.
|
|
64
|
+
|
|
65
|
+
Attributes:
|
|
66
|
+
k: Number of cards.
|
|
67
|
+
perms: `[k!, k]` permutations of the cards, in lexicographic order.
|
|
68
|
+
The first `m` cards of permutation `i` have rank `i // (k - m)!`
|
|
69
|
+
among partial permutations of length `m`.
|
|
70
|
+
pair_of, pair_of_base: Maps of every turn `t`, concatenated, and where
|
|
71
|
+
each starts. A map takes `rank(bids) * 3**t + winners` to the pair,
|
|
72
|
+
or -1 if no opponent bids produce those winners. Winners are base-3
|
|
73
|
+
digits, one per turn: 0 lost, 1 tied, 2 won.
|
|
74
|
+
last_pair_of: The map of the last decision turn, as a numpy array.
|
|
75
|
+
pair_bids, pair_opponent, pair_base: Per pair of every turn,
|
|
76
|
+
concatenated, the seat's bids and opponent bids that produce its
|
|
77
|
+
winners, padded with -1, and where each turn's pairs start.
|
|
78
|
+
num_pairs, offset: Per turn, the number of pairs and the number of
|
|
79
|
+
information sets before the turn's.
|
|
80
|
+
num_infosets: Number of information sets of a seat.
|
|
81
|
+
"""
|
|
82
|
+
|
|
83
|
+
def __init__(self, k):
|
|
84
|
+
self.k = k
|
|
85
|
+
perms = np.array(list(itertools.permutations(range(k))))
|
|
86
|
+
pair_of, pair_bids, pair_opponent, num_pairs = [], [], [], []
|
|
87
|
+
for t in range(k - 1): # The last turn plays itself.
|
|
88
|
+
mine = perms[:: math.factorial(k - t), :t]
|
|
89
|
+
witness = {}
|
|
90
|
+
for rank, bids in enumerate(mine):
|
|
91
|
+
for theirs in mine:
|
|
92
|
+
key = rank * 3**t + _code(np.sign(bids - theirs) + 1)
|
|
93
|
+
witness.setdefault(int(key), theirs)
|
|
94
|
+
keys = sorted(witness)
|
|
95
|
+
table = np.full(len(mine) * 3**t, -1)
|
|
96
|
+
table[keys] = np.arange(len(keys))
|
|
97
|
+
pair_of.append(table)
|
|
98
|
+
for key in keys:
|
|
99
|
+
pair_bids.append(
|
|
100
|
+
np.pad(mine[key // 3**t], (0, k - t), constant_values=-1)
|
|
101
|
+
)
|
|
102
|
+
pair_opponent.append(
|
|
103
|
+
np.pad(witness[key], (0, k - t), constant_values=-1)
|
|
104
|
+
)
|
|
105
|
+
num_pairs.append(len(keys))
|
|
106
|
+
self.offset = np.cumsum(
|
|
107
|
+
[0] + [math.perm(k, t + 1) * n for t, n in enumerate(num_pairs)]
|
|
108
|
+
)
|
|
109
|
+
self.num_infosets = int(self.offset[-1])
|
|
110
|
+
self.num_pairs = num_pairs
|
|
111
|
+
self.perms = jnp.asarray(perms)
|
|
112
|
+
self.last_pair_of = pair_of[-1]
|
|
113
|
+
self.pair_of = jnp.asarray(np.concatenate(pair_of))
|
|
114
|
+
self.pair_of_base = jnp.asarray(
|
|
115
|
+
np.cumsum([0] + [p.size for p in pair_of])
|
|
116
|
+
)
|
|
117
|
+
self.pair_base = jnp.asarray(np.cumsum([0, *num_pairs]))
|
|
118
|
+
self.pair_bids = jnp.asarray(np.array(pair_bids))
|
|
119
|
+
self.pair_opponent = jnp.asarray(np.array(pair_opponent))
|
|
120
|
+
|
|
121
|
+
def describe(self, ids):
|
|
122
|
+
"""Returns what identifies information sets, and a way to reach them.
|
|
123
|
+
|
|
124
|
+
Returns each information set's seat, turn, point-card rank, point
|
|
125
|
+
cards, bids, and opponent bids that produce its winners.
|
|
126
|
+
"""
|
|
127
|
+
k, offset = self.k, jnp.asarray(self.offset)
|
|
128
|
+
seat, local = jnp.divmod(ids, self.num_infosets)
|
|
129
|
+
turn = jnp.searchsorted(offset, local, side="right") - 1
|
|
130
|
+
num_pairs = jnp.asarray(self.num_pairs)[turn]
|
|
131
|
+
rank, pair = jnp.divmod(local - offset[turn], num_pairs)
|
|
132
|
+
pair = self.pair_base[turn] + pair
|
|
133
|
+
# Completing the point cards in increasing order gives a permutation.
|
|
134
|
+
factorial = jnp.asarray([math.factorial(n) for n in range(k + 1)])
|
|
135
|
+
points = self.perms[rank * factorial[k - turn - 1]]
|
|
136
|
+
return (
|
|
137
|
+
seat,
|
|
138
|
+
turn,
|
|
139
|
+
rank,
|
|
140
|
+
points,
|
|
141
|
+
self.pair_bids[pair],
|
|
142
|
+
self.pair_opponent[pair],
|
|
143
|
+
)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _code(digits):
|
|
147
|
+
"""Returns base-3 numbers of rows of digits, first digit most significant."""
|
|
148
|
+
return (digits * 3 ** np.arange(digits.shape[-1])[::-1]).sum(-1)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
@functools.partial(jax.jit, static_argnums=0)
|
|
152
|
+
def _parent_and_legal(tables, ids):
|
|
153
|
+
"""Returns the parent sequence and legal actions of information sets."""
|
|
154
|
+
k = tables.k
|
|
155
|
+
seat, turn, rank, _, bids, opponent = tables.describe(ids)
|
|
156
|
+
# The parent drops the last turn: its point card, bid, and winner.
|
|
157
|
+
previous = jnp.maximum(turn - 1, 0)
|
|
158
|
+
active = jnp.arange(k) < previous[:, None]
|
|
159
|
+
own_rank = jnp.zeros_like(ids)
|
|
160
|
+
for i in range(k):
|
|
161
|
+
smaller_unused = (jnp.arange(k) < bids[:, i : i + 1]).sum(1) - (
|
|
162
|
+
(bids[:, :i] < bids[:, i : i + 1]).sum(1)
|
|
163
|
+
)
|
|
164
|
+
own_rank = jnp.where(
|
|
165
|
+
active[:, i], own_rank * (k - i) + smaller_unused, own_rank
|
|
166
|
+
)
|
|
167
|
+
digits = jnp.where(active, jnp.sign(bids - opponent) + 1, 0)
|
|
168
|
+
power = jnp.maximum(previous[:, None] - 1 - jnp.arange(k), 0)
|
|
169
|
+
winners = (digits * 3**power).sum(1)
|
|
170
|
+
pair = tables.pair_of[
|
|
171
|
+
tables.pair_of_base[previous] + own_rank * 3**previous + winners
|
|
172
|
+
]
|
|
173
|
+
parent_rank = rank // (k - turn)
|
|
174
|
+
offset, num_pairs = (
|
|
175
|
+
jnp.asarray(tables.offset),
|
|
176
|
+
jnp.asarray(tables.num_pairs),
|
|
177
|
+
)
|
|
178
|
+
parent = (
|
|
179
|
+
seat * tables.num_infosets
|
|
180
|
+
+ offset[previous]
|
|
181
|
+
+ parent_rank * num_pairs[previous]
|
|
182
|
+
+ pair
|
|
183
|
+
)
|
|
184
|
+
last_bid = jnp.take_along_axis(bids, previous[:, None], 1)[:, 0]
|
|
185
|
+
parent = jnp.where(turn == 0, seat, k * (parent + 1) + last_bid)
|
|
186
|
+
legal = ~(bids[:, :, None] == jnp.arange(k)).any(1)
|
|
187
|
+
return parent.astype(jnp.int32), legal
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def _gradient(tables, plan):
|
|
191
|
+
"""Returns the gradient of `plan`, summing over all terminal histories."""
|
|
192
|
+
k, perms = tables.k, tables.perms
|
|
193
|
+
t = k - 2 # Last decision turn.
|
|
194
|
+
# For every pair of bid permutations: each seat's pair and last decision,
|
|
195
|
+
# and who won each turn (+1 seat 0, -1 seat 1).
|
|
196
|
+
i0, i1 = jnp.divmod(jnp.arange(perms.shape[0] ** 2), perms.shape[0])
|
|
197
|
+
b0, b1 = perms[i0], perms[i1]
|
|
198
|
+
code = 3 ** jnp.arange(t)[::-1]
|
|
199
|
+
last_pair_of = jnp.asarray(tables.last_pair_of)
|
|
200
|
+
pair0 = last_pair_of[
|
|
201
|
+
(i0 // 2) * 3**t + ((jnp.sign(b0 - b1) + 1)[:, :t] * code).sum(1)
|
|
202
|
+
]
|
|
203
|
+
pair1 = last_pair_of[
|
|
204
|
+
(i1 // 2) * 3**t + ((jnp.sign(b1 - b0) + 1)[:, :t] * code).sum(1)
|
|
205
|
+
]
|
|
206
|
+
won = jnp.sign(b0 - b1)
|
|
207
|
+
offset, num_pairs = int(tables.offset[t]), tables.num_pairs[t]
|
|
208
|
+
chunk = math.gcd(8, perms.shape[0]) # Point-card permutations per step.
|
|
209
|
+
|
|
210
|
+
def body(i, gradient):
|
|
211
|
+
points = jax.lax.dynamic_slice_in_dim(perms, i * chunk, chunk) + 1
|
|
212
|
+
weight = jnp.sign(won @ points.T) / perms.shape[0] # [pairs, chunk]
|
|
213
|
+
infoset = offset + (i * chunk + jnp.arange(chunk)) * num_pairs
|
|
214
|
+
seq0 = k * (infoset + pair0[:, None] + 1) + b0[:, t, None]
|
|
215
|
+
seq1 = (
|
|
216
|
+
k * (tables.num_infosets + infoset + pair1[:, None] + 1)
|
|
217
|
+
+ b1[:, t, None]
|
|
218
|
+
)
|
|
219
|
+
gradient = gradient.at[seq0].add(weight * plan[seq1])
|
|
220
|
+
return gradient.at[seq1].add(-weight * plan[seq0])
|
|
221
|
+
|
|
222
|
+
return jax.lax.fori_loop(
|
|
223
|
+
0, perms.shape[0] // chunk, body, jnp.zeros_like(plan)
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def _observe(game, tables, ids):
|
|
228
|
+
"""Returns observations of information sets, from representative states."""
|
|
229
|
+
seat, turn, _, points, bids, opponent = tables.describe(ids)
|
|
230
|
+
states = jax.tree.map(
|
|
231
|
+
lambda x: jnp.repeat(x[:1], ids.shape[0], 0), game.initial_states()[0]
|
|
232
|
+
)
|
|
233
|
+
own = jnp.arange(2)[None, :, None] == seat[:, None, None]
|
|
234
|
+
states = dataclasses.replace(
|
|
235
|
+
states,
|
|
236
|
+
point_cards=points,
|
|
237
|
+
bids=jnp.where(own, bids[:, None], opponent[:, None]),
|
|
238
|
+
turn=turn,
|
|
239
|
+
)
|
|
240
|
+
return jax.vmap(game.observe)(states, seat)
|