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