coggrid 0.2.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.
- coggrid/__init__.py +81 -0
- coggrid/config.py +240 -0
- coggrid/env.py +287 -0
- coggrid/generative.py +435 -0
- coggrid/observers.py +347 -0
- coggrid/viz/__init__.py +102 -0
- coggrid/viz/animate.py +1000 -0
- coggrid/viz/plots.py +1603 -0
- coggrid/viz/style.py +80 -0
- coggrid/world.py +378 -0
- coggrid-0.2.0.dist-info/METADATA +511 -0
- coggrid-0.2.0.dist-info/RECORD +14 -0
- coggrid-0.2.0.dist-info/WHEEL +4 -0
- coggrid-0.2.0.dist-info/licenses/LICENSE +21 -0
coggrid/__init__.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
"""coggrid — a stationary POMDP for studying compositional generalization
|
|
2
|
+
in latent space.
|
|
3
|
+
|
|
4
|
+
The "grid" is the table of observation rates over the joint realizations of the
|
|
5
|
+
active latent variables.
|
|
6
|
+
|
|
7
|
+
Quick start
|
|
8
|
+
-----------
|
|
9
|
+
>>> from coggrid import CogGridConfig, World, run_observers
|
|
10
|
+
>>> world = World(CogGridConfig(n_vars=200, n_contexts=2, seed=0))
|
|
11
|
+
>>> batch = world.sample_episodes(256)
|
|
12
|
+
>>> traces = run_observers(batch)
|
|
13
|
+
>>> traces["joint"].final()["accuracy"] > traces["naive"].final()["accuracy"]
|
|
14
|
+
True
|
|
15
|
+
|
|
16
|
+
Gymnasium interface
|
|
17
|
+
-------------------
|
|
18
|
+
>>> from coggrid import CogGridEnv
|
|
19
|
+
>>> env = CogGridEnv(seed=0)
|
|
20
|
+
>>> obs, info = env.reset()
|
|
21
|
+
>>> obs, reward, terminated, truncated, info = env.step(0)
|
|
22
|
+
|
|
23
|
+
Layout
|
|
24
|
+
------
|
|
25
|
+
``config`` :class:`CogGridConfig` — every tunable, validated, immutable.
|
|
26
|
+
``generative`` The generative model as pure functions.
|
|
27
|
+
``world`` :class:`World`, :class:`EpisodeBatch`, and the five swappable
|
|
28
|
+
generative-stage signatures.
|
|
29
|
+
``observers`` Ideal-observer baselines and metrics.
|
|
30
|
+
``env`` :class:`CogGridEnv`, the Gymnasium environment.
|
|
31
|
+
``viz`` Plotting. Every function returns a figure; none call ``show()``.
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
from __future__ import annotations
|
|
35
|
+
|
|
36
|
+
from .config import CogGridConfig
|
|
37
|
+
from .env import CogGridEnv
|
|
38
|
+
from .observers import (
|
|
39
|
+
BeliefTrace,
|
|
40
|
+
disentanglement,
|
|
41
|
+
factorization_regret,
|
|
42
|
+
joint_observer,
|
|
43
|
+
naive_observer,
|
|
44
|
+
run_observers,
|
|
45
|
+
score_belief,
|
|
46
|
+
)
|
|
47
|
+
from .world import (
|
|
48
|
+
ContextSampler,
|
|
49
|
+
EmbeddingSource,
|
|
50
|
+
EpisodeBatch,
|
|
51
|
+
LikelihoodModel,
|
|
52
|
+
ObservationModel,
|
|
53
|
+
RealizationSampler,
|
|
54
|
+
World,
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
__version__ = "0.2.0"
|
|
58
|
+
|
|
59
|
+
__all__ = [
|
|
60
|
+
"__version__",
|
|
61
|
+
# configuration
|
|
62
|
+
"CogGridConfig",
|
|
63
|
+
# environment
|
|
64
|
+
"World",
|
|
65
|
+
"EpisodeBatch",
|
|
66
|
+
"CogGridEnv",
|
|
67
|
+
# observers
|
|
68
|
+
"BeliefTrace",
|
|
69
|
+
"joint_observer",
|
|
70
|
+
"naive_observer",
|
|
71
|
+
"run_observers",
|
|
72
|
+
"score_belief",
|
|
73
|
+
"factorization_regret",
|
|
74
|
+
"disentanglement",
|
|
75
|
+
# extension points
|
|
76
|
+
"EmbeddingSource",
|
|
77
|
+
"ContextSampler",
|
|
78
|
+
"LikelihoodModel",
|
|
79
|
+
"RealizationSampler",
|
|
80
|
+
"ObservationModel",
|
|
81
|
+
]
|
coggrid/config.py
ADDED
|
@@ -0,0 +1,240 @@
|
|
|
1
|
+
"""Every tunable of a CogGrid world, in one validated, immutable dataclass.
|
|
2
|
+
|
|
3
|
+
Nothing in this module touches global state, allocates large arrays, or runs a
|
|
4
|
+
simulation.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import warnings
|
|
10
|
+
from dataclasses import dataclass, replace
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
import numpy as np
|
|
14
|
+
|
|
15
|
+
__all__ = ["CogGridConfig"]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(frozen=True, slots=True)
|
|
19
|
+
class CogGridConfig:
|
|
20
|
+
"""Immutable description of a CogGrid world.
|
|
21
|
+
|
|
22
|
+
The world is a *stationary* POMDP. On each episode:
|
|
23
|
+
|
|
24
|
+
1. ``n_contexts`` latent variables are drawn from a pool of ``n_vars``;
|
|
25
|
+
2. each active variable takes one of ``n_realizations`` discrete values;
|
|
26
|
+
3. those values jointly determine a Bernoulli rate for each of
|
|
27
|
+
``n_observations`` binary observation channels;
|
|
28
|
+
4. the agent sees ``n_steps`` i.i.d. samples of those channels and must
|
|
29
|
+
infer the value of the single *goal* variable.
|
|
30
|
+
|
|
31
|
+
Because the Bernoulli rates come from **pairwise** interactions between
|
|
32
|
+
active variables, the joint observation distribution does not factorise over
|
|
33
|
+
contexts. The gap between an observer that models the interactions and one
|
|
34
|
+
that assumes independence is the quantity of interest.
|
|
35
|
+
|
|
36
|
+
Attributes
|
|
37
|
+
----------
|
|
38
|
+
n_vars:
|
|
39
|
+
Size of the latent variable pool that contexts are drawn from.
|
|
40
|
+
n_contexts:
|
|
41
|
+
Number of simultaneously active latent variables per episode. Note that
|
|
42
|
+
the joint likelihood table is ``n_realizations ** n_contexts`` wide, so
|
|
43
|
+
this is the parameter that drives memory use.
|
|
44
|
+
n_realizations:
|
|
45
|
+
Number of discrete values each active variable can take.
|
|
46
|
+
n_observations:
|
|
47
|
+
Number of binary observation channels.
|
|
48
|
+
n_steps:
|
|
49
|
+
Observation samples per episode (the episode horizon).
|
|
50
|
+
embedding_dim:
|
|
51
|
+
Dimensionality of the per-variable key/query embeddings. Must be at least
|
|
52
|
+
``n_observations``, since embeddings are orthogonalized across channels.
|
|
53
|
+
likelihood_temp:
|
|
54
|
+
Scales the interaction potentials before the sigmoid. Larger values push
|
|
55
|
+
Bernoulli rates towards 0/1, making single observations more informative.
|
|
56
|
+
likelihood_freq:
|
|
57
|
+
Number of periods in the sinusoidal value profile. Higher values make
|
|
58
|
+
the mapping from interaction strength to realization multimodal.
|
|
59
|
+
batch_size:
|
|
60
|
+
Default batch size for :meth:`~coggrid.World.sample_episodes`.
|
|
61
|
+
n_held_out_vars:
|
|
62
|
+
Size of the held-out ("test") slice of the variable pool. ``None`` means
|
|
63
|
+
``n_vars // 10``, resolved to an int during validation. Held-out
|
|
64
|
+
variables are ``range(n_held_out_vars)``; training variables are the rest.
|
|
65
|
+
subsample_vars:
|
|
66
|
+
If set, draw contexts from a random subset of this size within each
|
|
67
|
+
split. Useful for probing generalization with a fixed small support.
|
|
68
|
+
allow_repeated_vars:
|
|
69
|
+
Whether a single episode may activate the same latent variable twice. Set
|
|
70
|
+
to ``False`` for distinct-variable episodes.
|
|
71
|
+
seed:
|
|
72
|
+
Seed for the default RNG. ``None`` means non-reproducible.
|
|
73
|
+
"""
|
|
74
|
+
|
|
75
|
+
n_vars: int = 500
|
|
76
|
+
n_contexts: int = 2
|
|
77
|
+
n_realizations: int = 10
|
|
78
|
+
n_observations: int = 5
|
|
79
|
+
n_steps: int = 30
|
|
80
|
+
embedding_dim: int = 30
|
|
81
|
+
likelihood_temp: float = 2.0
|
|
82
|
+
likelihood_freq: float = 1.0
|
|
83
|
+
batch_size: int = 1000
|
|
84
|
+
n_held_out_vars: int | None = None
|
|
85
|
+
subsample_vars: int | None = None
|
|
86
|
+
allow_repeated_vars: bool = True
|
|
87
|
+
seed: int | None = None
|
|
88
|
+
|
|
89
|
+
# ------------------------------------------------------------------ setup
|
|
90
|
+
def __post_init__(self) -> None:
|
|
91
|
+
positive = {
|
|
92
|
+
"n_vars": self.n_vars,
|
|
93
|
+
"n_contexts": self.n_contexts,
|
|
94
|
+
"n_realizations": self.n_realizations,
|
|
95
|
+
"n_observations": self.n_observations,
|
|
96
|
+
"n_steps": self.n_steps,
|
|
97
|
+
"embedding_dim": self.embedding_dim,
|
|
98
|
+
"batch_size": self.batch_size,
|
|
99
|
+
}
|
|
100
|
+
for name, value in positive.items():
|
|
101
|
+
if not isinstance(value, (int, np.integer)) or value < 1:
|
|
102
|
+
raise ValueError(f"{name} must be a positive int, got {value!r}")
|
|
103
|
+
|
|
104
|
+
if self.embedding_dim < self.n_observations:
|
|
105
|
+
raise ValueError(
|
|
106
|
+
f"embedding_dim ({self.embedding_dim}) must be >= n_observations "
|
|
107
|
+
f"({self.n_observations}): embeddings are orthogonalized across "
|
|
108
|
+
"observation channels, which is impossible in a lower-dimensional "
|
|
109
|
+
"space."
|
|
110
|
+
)
|
|
111
|
+
if self.likelihood_temp <= 0:
|
|
112
|
+
raise ValueError(f"likelihood_temp must be > 0, got {self.likelihood_temp}")
|
|
113
|
+
|
|
114
|
+
# Resolve the "auto" default once, so every reader downstream sees an int.
|
|
115
|
+
held_out = self.n_held_out_vars
|
|
116
|
+
if held_out is None:
|
|
117
|
+
held_out = self.n_vars // 10
|
|
118
|
+
object.__setattr__(self, "n_held_out_vars", held_out)
|
|
119
|
+
if not 0 <= held_out <= self.n_vars:
|
|
120
|
+
raise ValueError(
|
|
121
|
+
f"n_held_out_vars ({held_out}) must be in [0, n_vars={self.n_vars}]"
|
|
122
|
+
)
|
|
123
|
+
if held_out == 0:
|
|
124
|
+
raise ValueError(
|
|
125
|
+
"n_held_out_vars resolved to 0 — there would be no held-out "
|
|
126
|
+
"variables to evaluate generalization on. Increase n_vars "
|
|
127
|
+
"(>= 10) or set n_held_out_vars explicitly."
|
|
128
|
+
)
|
|
129
|
+
if held_out == self.n_vars:
|
|
130
|
+
raise ValueError(
|
|
131
|
+
"n_held_out_vars equals n_vars — there would be no training "
|
|
132
|
+
"variables left."
|
|
133
|
+
)
|
|
134
|
+
|
|
135
|
+
if self.subsample_vars is not None:
|
|
136
|
+
smallest = min(held_out, self.n_vars - held_out)
|
|
137
|
+
if not 1 <= self.subsample_vars <= smallest:
|
|
138
|
+
raise ValueError(
|
|
139
|
+
f"subsample_vars ({self.subsample_vars}) must be in "
|
|
140
|
+
f"[1, {smallest}] (the smaller of the two splits)"
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
if not self.allow_repeated_vars and self.n_contexts > self._split_floor():
|
|
144
|
+
raise ValueError(
|
|
145
|
+
f"n_contexts ({self.n_contexts}) exceeds the number of distinct "
|
|
146
|
+
f"variables available in the smaller split ({self._split_floor()}), "
|
|
147
|
+
"so distinct-variable episodes are impossible. Either increase "
|
|
148
|
+
"n_vars / n_held_out_vars or set allow_repeated_vars=True."
|
|
149
|
+
)
|
|
150
|
+
|
|
151
|
+
def _split_floor(self) -> int:
|
|
152
|
+
"""Distinct variables available in the smaller of the two splits."""
|
|
153
|
+
if self.subsample_vars is not None:
|
|
154
|
+
return self.subsample_vars
|
|
155
|
+
return min(self.n_held_out_vars, self.n_vars - self.n_held_out_vars)
|
|
156
|
+
|
|
157
|
+
# ------------------------------------------------------- derived quantities
|
|
158
|
+
@property
|
|
159
|
+
def n_train_vars(self) -> int:
|
|
160
|
+
"""Number of latent variables reserved for training episodes."""
|
|
161
|
+
return self.n_vars - self.n_held_out_vars
|
|
162
|
+
|
|
163
|
+
@property
|
|
164
|
+
def n_roll(self) -> int:
|
|
165
|
+
"""Length of the circular value profile (``1 + 2 * n_realizations``).
|
|
166
|
+
|
|
167
|
+
The profile is longer than ``n_realizations`` so that interaction
|
|
168
|
+
strengths can push probability mass "off the end" of the realization
|
|
169
|
+
axis and wrap around, rather than piling up at the boundary.
|
|
170
|
+
"""
|
|
171
|
+
return 1 + 2 * self.n_realizations
|
|
172
|
+
|
|
173
|
+
@property
|
|
174
|
+
def realization_shape(self) -> tuple[int, ...]:
|
|
175
|
+
"""Shape of the joint realization axes: ``(n_realizations,) * n_contexts``."""
|
|
176
|
+
return (self.n_realizations,) * self.n_contexts
|
|
177
|
+
|
|
178
|
+
def joint_likelihood_shape(self, batch_size: int | None = None) -> tuple[int, ...]:
|
|
179
|
+
"""Shape of the joint likelihood table for a batch."""
|
|
180
|
+
n = self.batch_size if batch_size is None else batch_size
|
|
181
|
+
return (n, self.n_observations, *self.realization_shape)
|
|
182
|
+
|
|
183
|
+
def memory_report(self, batch_size: int | None = None) -> str:
|
|
184
|
+
"""Human-readable estimate of the dominant allocation.
|
|
185
|
+
|
|
186
|
+
``n_realizations ** n_contexts`` grows fast; this is the number people
|
|
187
|
+
need in front of them *before* they wait on an OOM.
|
|
188
|
+
|
|
189
|
+
Examples
|
|
190
|
+
--------
|
|
191
|
+
>>> print(CogGridConfig(n_contexts=4, n_realizations=20).memory_report(1000))
|
|
192
|
+
1000 episodes x 5 channels x 20^4 realizations
|
|
193
|
+
joint likelihood : 6.0 GiB
|
|
194
|
+
joint belief : 35.8 GiB
|
|
195
|
+
peak (approx) : 41.7 GiB
|
|
196
|
+
"""
|
|
197
|
+
n = self.batch_size if batch_size is None else batch_size
|
|
198
|
+
table = int(np.prod(self.joint_likelihood_shape(n), dtype=np.int64)) * 8
|
|
199
|
+
belief = table * self.n_steps / self.n_observations
|
|
200
|
+
|
|
201
|
+
def human(x: float) -> str:
|
|
202
|
+
for unit in ("B", "KiB", "MiB", "GiB"):
|
|
203
|
+
if x < 1024:
|
|
204
|
+
return f"{x:.1f} {unit}"
|
|
205
|
+
x /= 1024
|
|
206
|
+
return f"{x:.1f} TiB"
|
|
207
|
+
|
|
208
|
+
return (
|
|
209
|
+
f"{n} episodes x {self.n_observations} channels x "
|
|
210
|
+
f"{self.n_realizations}^{self.n_contexts} realizations\n"
|
|
211
|
+
f" joint likelihood : {human(table)}\n"
|
|
212
|
+
f" joint belief : {human(belief)}\n"
|
|
213
|
+
f" peak (approx) : {human(table + belief)}"
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
def warn_if_large(
|
|
217
|
+
self, batch_size: int | None = None, limit_gib: float = 2.0
|
|
218
|
+
) -> None:
|
|
219
|
+
"""Emit a ``ResourceWarning`` when the joint table is likely to hurt."""
|
|
220
|
+
shape = self.joint_likelihood_shape(batch_size)
|
|
221
|
+
if int(np.prod(shape, dtype=np.int64)) * 8 > limit_gib * 1024**3:
|
|
222
|
+
warnings.warn(
|
|
223
|
+
"CogGrid joint likelihood is large:\n"
|
|
224
|
+
+ self.memory_report(batch_size)
|
|
225
|
+
+ "\nReduce batch_size, n_contexts or n_realizations, or sample "
|
|
226
|
+
"in chunks.",
|
|
227
|
+
ResourceWarning,
|
|
228
|
+
stacklevel=3,
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
# ---------------------------------------------------------------- helpers
|
|
232
|
+
def rng(self, seed: int | np.random.Generator | None = None) -> np.random.Generator:
|
|
233
|
+
"""Build a ``Generator``, preferring an explicit ``seed`` over ``self.seed``."""
|
|
234
|
+
if isinstance(seed, np.random.Generator):
|
|
235
|
+
return seed
|
|
236
|
+
return np.random.default_rng(self.seed if seed is None else seed)
|
|
237
|
+
|
|
238
|
+
def replace(self, **changes: Any) -> CogGridConfig:
|
|
239
|
+
"""Return a copy with ``changes`` applied (validation re-runs)."""
|
|
240
|
+
return replace(self, **changes)
|
coggrid/env.py
ADDED
|
@@ -0,0 +1,287 @@
|
|
|
1
|
+
"""Gymnasium-style environments.
|
|
2
|
+
|
|
3
|
+
The environment presents one episode at a time:
|
|
4
|
+
|
|
5
|
+
* **Observation** — a dict with three keys:
|
|
6
|
+
|
|
7
|
+
``observation``
|
|
8
|
+
``(n_observations,)`` int8 binary vector, one fresh sample per step.
|
|
9
|
+
``active_vars``
|
|
10
|
+
``(n_contexts,)`` int indices of the latent variables active this episode.
|
|
11
|
+
Constant within an episode. An agent is expected to *learn* an embedding
|
|
12
|
+
per index, which is where "generalization from experience" comes in — at
|
|
13
|
+
test time these indices have never been seen.
|
|
14
|
+
``goal_context``
|
|
15
|
+
Which column of ``active_vars`` the agent is scored on. Constant within
|
|
16
|
+
an episode.
|
|
17
|
+
|
|
18
|
+
* **Action** — ``Discrete(n_realizations)``: the agent's current guess at the
|
|
19
|
+
goal variable's value. Actions do **not** influence the observation stream; the
|
|
20
|
+
world is stationary and the agent is a decoder, not a controller. Reporting a
|
|
21
|
+
guess every step is what makes the accuracy-versus-time curves in the paper
|
|
22
|
+
well defined.
|
|
23
|
+
|
|
24
|
+
* **Reward** — 1.0 for a correct guess. ``reward_mode="dense"`` (default) pays
|
|
25
|
+
out every step; ``"terminal"`` pays out only on the last step.
|
|
26
|
+
|
|
27
|
+
Episodes end with ``terminated=True`` rather than ``truncated=True``. The
|
|
28
|
+
horizon is not an externally imposed time limit — it *is* the task, namely "how
|
|
29
|
+
well can you decode from exactly ``n_steps`` samples of evidence" — so there is
|
|
30
|
+
no value left to bootstrap from.
|
|
31
|
+
|
|
32
|
+
Because the generative model is cheap in batch and expensive one episode at a
|
|
33
|
+
time, :class:`CogGridEnv` samples a buffer of episodes and serves them one by
|
|
34
|
+
one. If you want the batch directly, call
|
|
35
|
+
:meth:`~coggrid.World.sample_episodes` — that is the vectorized path, and no
|
|
36
|
+
environment wrapper is needed for it.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
from __future__ import annotations
|
|
40
|
+
|
|
41
|
+
from typing import Any, Literal
|
|
42
|
+
|
|
43
|
+
import numpy as np
|
|
44
|
+
from gymnasium import Env
|
|
45
|
+
from gymnasium.spaces import Box, Dict, Discrete, MultiBinary
|
|
46
|
+
|
|
47
|
+
from .config import CogGridConfig
|
|
48
|
+
from .generative import Split
|
|
49
|
+
from .world import EpisodeBatch, World
|
|
50
|
+
|
|
51
|
+
__all__ = ["CogGridEnv", "register"]
|
|
52
|
+
|
|
53
|
+
RewardMode = Literal["dense", "terminal"]
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _observation_space(cfg: CogGridConfig) -> Dict:
|
|
57
|
+
return Dict(
|
|
58
|
+
{
|
|
59
|
+
"observation": MultiBinary(cfg.n_observations),
|
|
60
|
+
"active_vars": Box(
|
|
61
|
+
low=0,
|
|
62
|
+
high=cfg.n_vars - 1,
|
|
63
|
+
shape=(cfg.n_contexts,),
|
|
64
|
+
dtype=np.int64,
|
|
65
|
+
),
|
|
66
|
+
"goal_context": Discrete(cfg.n_contexts),
|
|
67
|
+
}
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class CogGridEnv(Env):
|
|
72
|
+
"""Single-episode Gymnasium environment.
|
|
73
|
+
|
|
74
|
+
Parameters
|
|
75
|
+
----------
|
|
76
|
+
config:
|
|
77
|
+
A :class:`CogGridConfig`, or ``None`` for the defaults.
|
|
78
|
+
world:
|
|
79
|
+
An existing :class:`World` to draw from. Pass one to share embeddings
|
|
80
|
+
(and therefore the train/held-out split) across several environments.
|
|
81
|
+
If given, ``config`` is ignored.
|
|
82
|
+
split:
|
|
83
|
+
Which variable pool episodes are drawn from — ``"held_out"`` or ``"train"``.
|
|
84
|
+
reward_mode:
|
|
85
|
+
``"dense"`` or ``"terminal"``.
|
|
86
|
+
buffer_size:
|
|
87
|
+
How many episodes to generate per internal refill. Larger is faster but
|
|
88
|
+
uses more memory; see :meth:`CogGridConfig.memory_report`.
|
|
89
|
+
expose_likelihood:
|
|
90
|
+
If ``True``, ``info`` carries the full joint and marginal rate tables for
|
|
91
|
+
the current episode. This is what an ideal-observer baseline needs. Off
|
|
92
|
+
by default because the arrays are large and an agent should not see them.
|
|
93
|
+
seed:
|
|
94
|
+
Seed for the world's embeddings and the episode stream.
|
|
95
|
+
|
|
96
|
+
Examples
|
|
97
|
+
--------
|
|
98
|
+
>>> env = CogGridEnv(seed=0)
|
|
99
|
+
>>> obs, info = env.reset()
|
|
100
|
+
>>> sorted(obs)
|
|
101
|
+
['active_vars', 'goal_context', 'observation']
|
|
102
|
+
>>> total = 0.0
|
|
103
|
+
>>> for _ in range(env.cfg.n_steps):
|
|
104
|
+
... obs, reward, terminated, truncated, info = env.step(env.action_space.sample())
|
|
105
|
+
... total += reward
|
|
106
|
+
>>> terminated
|
|
107
|
+
True
|
|
108
|
+
"""
|
|
109
|
+
|
|
110
|
+
metadata = {"render_modes": ["ansi"]}
|
|
111
|
+
|
|
112
|
+
def __init__(
|
|
113
|
+
self,
|
|
114
|
+
config: CogGridConfig | None = None,
|
|
115
|
+
*,
|
|
116
|
+
world: World | None = None,
|
|
117
|
+
split: Split = "held_out",
|
|
118
|
+
reward_mode: RewardMode = "dense",
|
|
119
|
+
buffer_size: int = 256,
|
|
120
|
+
expose_likelihood: bool = False,
|
|
121
|
+
seed: int | None = None,
|
|
122
|
+
render_mode: str | None = None,
|
|
123
|
+
) -> None:
|
|
124
|
+
if reward_mode not in ("dense", "terminal"):
|
|
125
|
+
raise ValueError(
|
|
126
|
+
f"reward_mode must be 'dense' or 'terminal', got {reward_mode!r}"
|
|
127
|
+
)
|
|
128
|
+
if buffer_size < 1:
|
|
129
|
+
raise ValueError(f"buffer_size must be >= 1, got {buffer_size}")
|
|
130
|
+
|
|
131
|
+
if world is None:
|
|
132
|
+
cfg = config if config is not None else CogGridConfig()
|
|
133
|
+
if seed is not None:
|
|
134
|
+
cfg = cfg.replace(seed=seed)
|
|
135
|
+
world = World(cfg)
|
|
136
|
+
self.world = world
|
|
137
|
+
self.cfg = world.cfg
|
|
138
|
+
|
|
139
|
+
self.split: Split = split
|
|
140
|
+
self.reward_mode: RewardMode = reward_mode
|
|
141
|
+
self.buffer_size = int(buffer_size)
|
|
142
|
+
self.expose_likelihood = bool(expose_likelihood)
|
|
143
|
+
self.render_mode = render_mode
|
|
144
|
+
|
|
145
|
+
self.observation_space = _observation_space(self.cfg)
|
|
146
|
+
self.action_space = Discrete(self.cfg.n_realizations)
|
|
147
|
+
|
|
148
|
+
self._stream = np.random.default_rng(
|
|
149
|
+
self.cfg.seed if seed is None else seed
|
|
150
|
+
)
|
|
151
|
+
self._buffer: EpisodeBatch | None = None
|
|
152
|
+
self._cursor = 0
|
|
153
|
+
self._episode: EpisodeBatch | None = None
|
|
154
|
+
self._t = 0
|
|
155
|
+
self._last_action: int | None = None
|
|
156
|
+
|
|
157
|
+
# ------------------------------------------------------------------ buffer
|
|
158
|
+
def _next_episode(self) -> EpisodeBatch:
|
|
159
|
+
if self._buffer is None or self._cursor >= len(self._buffer):
|
|
160
|
+
self._buffer = self.world.sample_episodes(
|
|
161
|
+
self.buffer_size, split=self.split, rng=self._stream
|
|
162
|
+
)
|
|
163
|
+
self._cursor = 0
|
|
164
|
+
episode = self._buffer.select(self._cursor)
|
|
165
|
+
self._cursor += 1
|
|
166
|
+
return episode
|
|
167
|
+
|
|
168
|
+
# ---------------------------------------------------------------- gym api
|
|
169
|
+
def reset(
|
|
170
|
+
self,
|
|
171
|
+
*,
|
|
172
|
+
seed: int | None = None,
|
|
173
|
+
options: dict[str, Any] | None = None,
|
|
174
|
+
) -> tuple[dict[str, Any], dict[str, Any]]:
|
|
175
|
+
"""Start a new episode. Returns ``(observation, info)``."""
|
|
176
|
+
if seed is not None:
|
|
177
|
+
self._stream = np.random.default_rng(seed)
|
|
178
|
+
self._buffer = None
|
|
179
|
+
if options:
|
|
180
|
+
split = options.get("split")
|
|
181
|
+
if split is not None:
|
|
182
|
+
self.split = split
|
|
183
|
+
self._buffer = None
|
|
184
|
+
|
|
185
|
+
self._episode = self._next_episode()
|
|
186
|
+
self._t = 0
|
|
187
|
+
self._last_action = None
|
|
188
|
+
return self._observe(), self._info()
|
|
189
|
+
|
|
190
|
+
def step(
|
|
191
|
+
self, action: int | np.integer
|
|
192
|
+
) -> tuple[dict[str, Any], float, bool, bool, dict[str, Any]]:
|
|
193
|
+
"""Submit a guess at the goal variable's value and advance one step."""
|
|
194
|
+
if self._episode is None:
|
|
195
|
+
raise RuntimeError("call reset() before step()")
|
|
196
|
+
|
|
197
|
+
action = int(action)
|
|
198
|
+
if not 0 <= action < self.cfg.n_realizations:
|
|
199
|
+
raise ValueError(
|
|
200
|
+
f"action {action} outside Discrete({self.cfg.n_realizations})"
|
|
201
|
+
)
|
|
202
|
+
self._last_action = action
|
|
203
|
+
|
|
204
|
+
correct = action == int(self._episode.goal_value[0])
|
|
205
|
+
self._t += 1
|
|
206
|
+
terminated = self._t >= self.cfg.n_steps
|
|
207
|
+
|
|
208
|
+
if self.reward_mode == "dense":
|
|
209
|
+
reward = float(correct)
|
|
210
|
+
else:
|
|
211
|
+
reward = float(correct) if terminated else 0.0
|
|
212
|
+
|
|
213
|
+
# On the final step there is no further observation to reveal; repeat the
|
|
214
|
+
# last one so the returned dict always matches observation_space.
|
|
215
|
+
observation = self._observe(clamp=True)
|
|
216
|
+
info = self._info()
|
|
217
|
+
info["correct"] = correct
|
|
218
|
+
return observation, reward, terminated, False, info
|
|
219
|
+
|
|
220
|
+
def render(self) -> str | None:
|
|
221
|
+
"""One-line text summary of the current step."""
|
|
222
|
+
if self.render_mode != "ansi" or self._episode is None:
|
|
223
|
+
return None
|
|
224
|
+
e = self._episode
|
|
225
|
+
bits = "".join("1" if b else "0" for b in self._current_bits())
|
|
226
|
+
guess = "-" if self._last_action is None else str(self._last_action)
|
|
227
|
+
return (
|
|
228
|
+
f"t={self._t:>3}/{self.cfg.n_steps} obs={bits} "
|
|
229
|
+
f"vars={e.ctx_inds[0].tolist()} goal_ctx={int(e.goal_ind[0])} "
|
|
230
|
+
f"truth={int(e.goal_value[0])} guess={guess}"
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
# ---------------------------------------------------------------- internals
|
|
234
|
+
def _current_bits(self, clamp: bool = True) -> np.ndarray:
|
|
235
|
+
assert self._episode is not None
|
|
236
|
+
t = min(self._t, self.cfg.n_steps - 1) if clamp else self._t
|
|
237
|
+
return self._episode.observations[0, t]
|
|
238
|
+
|
|
239
|
+
def _observe(self, clamp: bool = True) -> dict[str, Any]:
|
|
240
|
+
assert self._episode is not None
|
|
241
|
+
return {
|
|
242
|
+
"observation": self._current_bits(clamp).astype(np.int8),
|
|
243
|
+
"active_vars": self._episode.ctx_inds[0].astype(np.int64),
|
|
244
|
+
"goal_context": int(self._episode.goal_ind[0]),
|
|
245
|
+
}
|
|
246
|
+
|
|
247
|
+
def _info(self) -> dict[str, Any]:
|
|
248
|
+
assert self._episode is not None
|
|
249
|
+
e = self._episode
|
|
250
|
+
info: dict[str, Any] = {
|
|
251
|
+
"step": self._t,
|
|
252
|
+
"split": self.split,
|
|
253
|
+
# Ground truth: for analysis and for scoring, never as agent input.
|
|
254
|
+
"goal_value": int(e.goal_value[0]),
|
|
255
|
+
"context_values": e.ctx_vals[0].copy(),
|
|
256
|
+
"true_rates": e.true_rates[0].copy(),
|
|
257
|
+
}
|
|
258
|
+
if self.expose_likelihood:
|
|
259
|
+
info["rates"] = e.rates[0]
|
|
260
|
+
info["marginal_rates"] = e.marginal_rates[0]
|
|
261
|
+
return info
|
|
262
|
+
|
|
263
|
+
# ------------------------------------------------------------------ extras
|
|
264
|
+
@property
|
|
265
|
+
def episode(self) -> EpisodeBatch | None:
|
|
266
|
+
"""The current episode as a one-element :class:`EpisodeBatch`.
|
|
267
|
+
|
|
268
|
+
Handy for running an ideal observer on exactly the episode the agent saw.
|
|
269
|
+
"""
|
|
270
|
+
return self._episode
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
def register() -> None:
|
|
274
|
+
"""Add ``CogGrid-v0`` to the Gymnasium registry.
|
|
275
|
+
|
|
276
|
+
>>> import gymnasium as gym
|
|
277
|
+
>>> from coggrid.env import register
|
|
278
|
+
>>> register()
|
|
279
|
+
>>> env = gym.make("CogGrid-v0")
|
|
280
|
+
"""
|
|
281
|
+
from gymnasium.envs.registration import register as gym_register
|
|
282
|
+
|
|
283
|
+
gym_register(
|
|
284
|
+
id="CogGrid-v0",
|
|
285
|
+
entry_point="coggrid.env:CogGridEnv",
|
|
286
|
+
max_episode_steps=None, # the env terminates on its own
|
|
287
|
+
)
|