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 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)