coggrid 0.2.0__tar.gz

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.
@@ -0,0 +1,13 @@
1
+ # Python
2
+ __pycache__/
3
+ *.py[cod]
4
+ .pytest_cache/
5
+ .ruff_cache/
6
+ .mypy_cache/
7
+ *.egg-info/
8
+ build/
9
+ dist/
10
+
11
+ # Environments
12
+ .venv/
13
+ venv/
coggrid-0.2.0/LICENSE ADDED
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 John Schwarcz
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
coggrid-0.2.0/PKG-INFO ADDED
@@ -0,0 +1,511 @@
1
+ Metadata-Version: 2.5
2
+ Name: coggrid
3
+ Version: 0.2.0
4
+ Summary: A stationary POMDP for studying compositional generalization in latent space
5
+ Project-URL: Homepage, https://github.com/johnschwarcz/coggrid
6
+ Project-URL: Repository, https://github.com/johnschwarcz/coggrid
7
+ Project-URL: Issues, https://github.com/johnschwarcz/coggrid/issues
8
+ Project-URL: Paper, https://arxiv.org/abs/2603.27134
9
+ Author-email: John Schwarcz <Johnschwarcz@gmail.com>
10
+ License-Expression: MIT
11
+ License-File: LICENSE
12
+ Keywords: bayesian-inference,compositional-generalization,disentanglement,gymnasium,ideal-observer,pomdp,reinforcement-learning,toy-model
13
+ Classifier: Development Status :: 4 - Beta
14
+ Classifier: Intended Audience :: Science/Research
15
+ Classifier: Operating System :: OS Independent
16
+ Classifier: Programming Language :: Python :: 3.10
17
+ Classifier: Programming Language :: Python :: 3.11
18
+ Classifier: Programming Language :: Python :: 3.12
19
+ Classifier: Programming Language :: Python :: 3.13
20
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
21
+ Requires-Python: >=3.10
22
+ Requires-Dist: gymnasium>=0.29
23
+ Requires-Dist: numpy>=1.24
24
+ Provides-Extra: dev
25
+ Requires-Dist: matplotlib>=3.7; extra == 'dev'
26
+ Requires-Dist: mypy>=1.8; extra == 'dev'
27
+ Requires-Dist: pytest>=7.4; extra == 'dev'
28
+ Requires-Dist: ruff>=0.5; extra == 'dev'
29
+ Requires-Dist: seaborn>=0.13; extra == 'dev'
30
+ Provides-Extra: viz
31
+ Requires-Dist: matplotlib>=3.7; extra == 'viz'
32
+ Requires-Dist: seaborn>=0.13; extra == 'viz'
33
+ Description-Content-Type: text/markdown
34
+
35
+ # coggrid
36
+
37
+ [![CI](https://github.com/johnschwarcz/coggrid/actions/workflows/ci.yml/badge.svg)](https://github.com/johnschwarcz/coggrid/actions/workflows/ci.yml)
38
+ [![PyPI](https://img.shields.io/pypi/v/coggrid)](https://pypi.org/project/coggrid/)
39
+ [![Python 3.10+](https://img.shields.io/badge/python-3.10%2B-blue)](https://www.python.org/downloads/)
40
+ [![License: MIT](https://img.shields.io/badge/license-MIT-green)](https://github.com/johnschwarcz/coggrid/blob/main/LICENSE)
41
+ [![arXiv](https://img.shields.io/badge/arXiv-2603.27134-b31b1b)](https://arxiv.org/abs/2603.27134)
42
+
43
+ A stationary POMDP for studying **compositional generalization in latent space** —
44
+ the environment and ideal-observer baselines from
45
+ [arXiv:2603.27134](https://arxiv.org/abs/2603.27134).
46
+
47
+ > **Sequential Bayesian inference is represented as navigating a latent space.**
48
+
49
+ An example episode:
50
+ <p align="center">
51
+ <img src="https://raw.githubusercontent.com/johnschwarcz/coggrid/main/docs/images/episode_animation_1.gif" alt="An episode playing back: joint and naive posteriors, their difference, and the evidence stream" width="100%">
52
+ </p>
53
+
54
+ <p align="center"><em>
55
+ Top: the optimal and naive observers' posteriors and their difference. Bottom: A stream of observations.
56
+ </em></p>
57
+
58
+
59
+ ---
60
+
61
+ ## Contrasting Optimal and Naive Bayes
62
+
63
+ <p align="center">
64
+ <img src="https://raw.githubusercontent.com/johnschwarcz/coggrid/main/docs/images/factorization_cost.png" alt="Factorization regret predicts the naive observer's failures, and its own confidence does not" width="100%">
65
+ </p>
66
+
67
+ The naive observer is not just noisier — Factorization regret decouples confidence and accuracy.
68
+
69
+ ---
70
+
71
+ ## How the environment works
72
+
73
+ ```python
74
+ from coggrid import CogGridConfig, World
75
+ from coggrid.viz import animate_interaction_phases, plot_interaction_phases
76
+
77
+ world = World(CogGridConfig(n_vars=500, n_contexts=2, seed=0))
78
+ batch = world.sample_episodes(2000)
79
+
80
+ plot_interaction_phases(world, batch, episode=0, channels=(0,))
81
+ animate_interaction_phases(world, batch, episode=0)
82
+ ```
83
+
84
+ Each episode:
85
+
86
+ 1. **Contexts.** `n_contexts` latent variables are drawn from a pool of `n_vars`.
87
+ One is designated the **goal**.
88
+ 2. **Realizations.** Each "active" variable in the context takes one of `n_realizations` discrete
89
+ values.
90
+ 3. **Likelihood.** Active variables interact through key/query embeddings which map to an `n_realizations`^`n_contexts` likelihood.
91
+ 5. **Observations.** `n_steps` i.i.d. observations sampled stochastically from the likelihood / prob. of observing 1 aka 1 - prob. of observing 0.
92
+
93
+ ### What's in a batch
94
+
95
+ `sample_episodes` returns an `EpisodeBatch`:
96
+
97
+ | Field | Shape | Meaning |
98
+ | --- | --- | --- |
99
+ | `cfg` | — | The `CogGridConfig` the episodes were drawn from. |
100
+ | `split` | — | `"train"` or `"held_out"` — which variable pool the contexts came from. |
101
+ | `ctx_inds` | `(batch_size, n_contexts)` | Indices of the active variables, sorted. |
102
+ | `goal_ind` | `(batch_size,)` | Which column of `ctx_inds` is the goal. |
103
+ | `ctx_vals` | `(batch_size, n_contexts)` | The realization each active variable took. |
104
+ | `goal_value` | `(batch_size,)` | The goal variable's realization — the quantity to infer. |
105
+ | `observations` | `(batch_size, n_steps, n_observations)` | The boolean observation stream. |
106
+ | `rates` | `(batch_size, n_observations, n_realizations, ...)` | The Bernoulli rate of an observation for every joint realization. |
107
+ | `marginal_rates` | `(batch_size, n_observations, n_contexts, n_realizations)` | The same for each variable's realizations after marginalizing over all others. |
108
+ | `true_rates` | `(batch_size, n_observations)` | `rates` at `ctx_vals` — the rates that actually generated `observations`. |
109
+ | `interactions` | `(batch_size, n_observations, k)` | Raw interaction terms, one per ordered pair of active variables. |
110
+
111
+ A learning agent should see only `ctx_inds`, `goal_ind` and `observations` (and optionally `interactions` if they are not to be learned).
112
+
113
+ `batch.select(index)` returns a sub-batch, `batch.save(path)` / `EpisodeBatch.load(path)` round-trip it through a compressed `.npz`.
114
+
115
+ ### Interactions shift the joint likelihood in latent space
116
+
117
+ The inner product of one variable's key with the other's query sets a *phase* that translates an "XOR"-like pattern.
118
+
119
+ <p align="center">
120
+ <img src="https://raw.githubusercontent.com/johnschwarcz/coggrid/main/docs/images/interaction_phases.png" alt="Variable embeddings, the inner products they produce, and the standard likelihood pattern those phases translate" width="100%">
121
+ </p>
122
+
123
+
124
+ Reading left to right:
125
+ 1. **Embeddings.** Each latent variable carries a **key** and a **query** vector
126
+ per channel. The angle between one variable's key and the other's query gives a score `z`.
127
+ 2. **Phase.** Each score shifts a standard sinusoid by `−2π · likelihood_freq · z`.
128
+ 3. **The standard pattern**, Sinasoids are expanded to a repeating pattern through an outer product.
129
+ 4. **The selected pattern** A specific pair of scores defines a specific region of the pattern.
130
+
131
+ ### Impact of rotating the embeddings
132
+ Turning the keys sweeps the scores, and every panel to the right follows:
133
+
134
+ <p align="center">
135
+ <img src="https://raw.githubusercontent.com/johnschwarcz/coggrid/main/docs/images/interaction_phases_animated.gif" alt="Turning a variable's key sweeps its interaction score, sliding the selected region across the standard pattern" width="100%">
136
+ </p>
137
+
138
+ Note: This is meant to demonstrate how embeddings impact the likelihood. Embeddings are never actually rotated.
139
+
140
+ ### What a single observation actually says
141
+
142
+ ```python
143
+ from coggrid.viz import plot_evidence_likelihood
144
+
145
+ plot_evidence_likelihood(batch, episode=0)
146
+ ```
147
+
148
+ The agent sees a **vector** of
149
+ `n_observations` at once, and its likelihood is the product of the
150
+ per-channel rates.
151
+
152
+ <p align="center">
153
+ <img src="https://raw.githubusercontent.com/johnschwarcz/coggrid/main/docs/images/evidence_likelihood.png" alt="Per-channel rate tables, and the posterior induced by every possible observation vector" width="100%">
154
+ </p>
155
+
156
+
157
+ The top row is the per-channel joint likelihoods.
158
+ The grid below is every possible observation vector and the belief update it induces.
159
+
160
+
161
+ ## Quickstart
162
+
163
+ ```python
164
+ from coggrid import CogGridConfig, World, run_observers
165
+
166
+ world = World(CogGridConfig(n_vars=500, n_contexts=2, seed=0))
167
+ batch = world.sample_episodes(2000)
168
+ traces = run_observers(batch)
169
+
170
+ traces["joint"].final() # {'accuracy': ..., 'p_correct': ..., 'mse': ...}
171
+ traces["naive"].final() # the same three, for the factorized observer
172
+ ```
173
+
174
+ ## Install
175
+
176
+ ```bash
177
+ pip install coggrid # core
178
+ pip install "coggrid[viz]" # + matplotlib and seaborn for the figures
179
+ ```
180
+
181
+ Python 3.10+. No compiled extensions, no GPU, any OS.
182
+
183
+ From source, for development:
184
+
185
+ ```bash
186
+ git clone https://github.com/johnschwarcz/coggrid
187
+ cd coggrid
188
+ pip install -e ".[dev]"
189
+ pytest
190
+ ```
191
+
192
+ ### Every knob
193
+
194
+ `CogGridConfig` holds the whole specification of a world.
195
+
196
+ | Field | Default | Meaning |
197
+ | --- | --- | --- |
198
+ | `n_vars` | 500 | Size of the latent-variable pool episodes draw from. |
199
+ | `n_contexts` | 2 | Variables active per episode. The joint table is `n_realizations ** n_contexts` wide, so this drives memory. |
200
+ | `n_realizations` | 10 | Discrete values each active variable can take. |
201
+ | `n_observations` | 5 | Binary observation channels. |
202
+ | `n_steps` | 30 | Observation samples per episode — the horizon. |
203
+ | `embedding_dim` | 30 | Length of each key/query vector. Must be ≥ `n_observations`, since they are orthogonalized across channels. |
204
+ | `likelihood_temp` | 2.0 | Scales the potentials before the sigmoid. Higher pushes rates towards 0/1, making single observations more informative. |
205
+ | `likelihood_freq` | 1.0 | Periods in the value profile. Higher partitions the realization axis more finely. |
206
+ | `batch_size` | 1000 | Default batch size for `sample_episodes`. |
207
+ | `n_held_out_vars` | `None` | Size of the held-out pool. `None` means `n_vars // 10`. Held-out variables are `range(n_held_out_vars)`. |
208
+ | `subsample_vars` | `None` | If set, draw contexts from a random subset of this size within each split. |
209
+ | `allow_repeated_vars` | `True` | Whether one episode may activate the same variable twice. |
210
+ | `seed` | `None` | Seed for the default RNG. `None` means a fresh world each run. |
211
+
212
+ ### The two baselines
213
+
214
+ | Observer | Uses | Meaning |
215
+ | --- | --- | --- |
216
+ | `joint` | `batch.rates` — the full interaction table | Bayes-optimal. The ceiling. |
217
+ | `naive` | `batch.marginal_rates` — interactions averaged away | Exactly the joint observer with the interactions removed, and nothing else changed. |
218
+
219
+ Both run in log space with a cumulative sum rather than a per-step
220
+ multiply-and-renormalize, so neither underflows at long horizons (tested to
221
+ `n_steps=4000`).
222
+
223
+ ```python
224
+ from coggrid import CogGridConfig, World, run_observers
225
+ from coggrid import factorization_regret, disentanglement
226
+
227
+ world = World(CogGridConfig(n_vars=500, n_contexts=2, seed=0))
228
+ batch = world.sample_episodes(2000)
229
+ traces = run_observers(batch)
230
+
231
+ factorization_regret(traces["joint"], traces["naive"]) # (batch_size, n_steps)
232
+ disentanglement(traces["joint"], traces["naive"], batch) # (batch_size, n_steps)
233
+ ```
234
+
235
+ **Factorization regret** — `D_KL(B_joint ‖ B_naive)`
236
+ ([§3.1](https://arxiv.org/abs/2603.27134)) — measures the cost of factorizing inference of interacting latent variables.
237
+
238
+ **Dis-entanglement** ([§B.3](https://arxiv.org/abs/2603.27134)) quantifies the 'history-dependence' of an observation,
239
+ by measuring the divergence between the naive marginal belief update `p(o_t | r)`
240
+ and the optimal marginal belief update `p(o_t | r, o_1:t-1)`.
241
+ It is zero when marginal belief dynamics are Markovian.
242
+
243
+
244
+ ### Train / held-out splits
245
+
246
+ ```python
247
+ world.sample_episodes(1000, split="train") # <=1 novel variable, goal always familiar
248
+ world.sample_episodes(1000, split="held_out") # every active variable novel
249
+ ```
250
+
251
+ `"train"` is the condition under which a subset of the variables are never the
252
+ goal and never co-occur. `"held_out"` is the condition where no variable in the
253
+ episode has been a goal, nor co-occurred, during training. This is specifically relevant for training and evaluating networks.
254
+
255
+ ---
256
+
257
+ ## Gymnasium interface
258
+
259
+ ```python
260
+ from coggrid import CogGridEnv
261
+
262
+ env = CogGridEnv(seed=0)
263
+ obs, info = env.reset()
264
+
265
+ for _ in range(env.cfg.n_steps):
266
+ obs, reward, terminated, truncated, info = env.step(my_agent(obs))
267
+ ```
268
+
269
+ **obs** — a dict:
270
+
271
+ | Key | Shape | Meaning |
272
+ | --- | --- | --- |
273
+ | `observation` | `(n_observations,)` int8 | Fresh binary sample each step. |
274
+ | `active_vars` | `(n_contexts,)` int64 | Which latent variables are active. Constant within an episode. |
275
+ | `goal_context` | scalar int | Which column of `active_vars` you are scored on. |
276
+
277
+ `active_vars` holds *indices*, not embeddings — learning an embedding per index
278
+ from experience is the point of the task. At test time those indices have never
279
+ been seen.
280
+
281
+ **Action** — `Discrete(n_realizations)`: the agent's current guess (MAP) of the goal
282
+ variable's value.
283
+
284
+ **Reward** — 1.0 for a correct guess, every step (`reward_mode="dense"`) or only
285
+ on the last step (`reward_mode="terminal"`).
286
+
287
+ Episodes end with `terminated=True`, not `truncated=True`.
288
+
289
+ `coggrid.env.register()` adds `CogGrid-v0` to the Gymnasium registry.
290
+
291
+ `CogGridEnv` takes:
292
+
293
+ | Argument | Default | Meaning |
294
+ | --- | --- | --- |
295
+ | `config` | `None` | The world specification, per the table above. `None` builds a default `CogGridConfig()`. |
296
+ | `world` | `None` | Draw from an existing `World` instead of building one, so several environments share embeddings and therefore the same split. Overrides `config`. |
297
+ | `split` | `"held_out"` | Which variable pool episodes come from — see below. |
298
+ | `reward_mode` | `"dense"` | `"dense"` scores every step, `"terminal"` only the last. |
299
+ | `buffer_size` | 256 | Episodes generated per internal refill. Sampling in blocks is much cheaper than one at a time; larger uses more memory. |
300
+ | `expose_likelihood` | `False` | Put the joint and marginal rate tables in `info` — what an ideal-observer baseline needs, and what a learning agent must not see. |
301
+ | `render_mode` | `None` | `"ansi"` makes `render()` return a text summary of the current step. |
302
+ | `seed` | `None` | Seeds the world's embeddings and the episode stream. `None` means fresh episodes each run. |
303
+
304
+ ### Batched rollouts
305
+
306
+ There is no vector environment. Batching is what `World` already does:
307
+
308
+ ```python
309
+ rollout = world.sample_episodes(4096) # every episode at once
310
+ rollout.observations # (4096, n_steps, n_observations)
311
+ ```
312
+
313
+ `sample_episodes` builds the whole batch with one vectorized einsum over the
314
+ joint likelihood, which is the cheap path. A gym wrapper around it would only add
315
+ a step-by-step interface on top of data that is already complete.
316
+
317
+ ---
318
+
319
+ ## Visualization
320
+
321
+ Every plotting function takes data, returns a `Figure`, and never calls `show()`.
322
+
323
+ ```python
324
+ from coggrid.viz import summary_figure
325
+
326
+ figures = summary_figure(batch, traces) # -> 4 figures, narrowing in scope
327
+ ```
328
+
329
+ - **`plot_episode`** — one episode's rate surface beside its belief traces (shown
330
+ above).
331
+ - **`plot_performance`** — batch-averaged curves, one column per metric
332
+ (accuracy, `p_correct`, MSE, factorization regret), with individual episode
333
+ trajectories behind each mean.
334
+ - **`plot_regret_analysis`** — whether regret explains the gap.
335
+ - **`plot_belief_shape`** — the posteriors behind the accuracy numbers: the mean
336
+ final belief with every episode aligned on its true value, and the
337
+ distribution of `P(true value)` across episodes at each timestep.
338
+
339
+ The last two are omitted when `n_contexts == 1`, where the two observers coincide
340
+ and there is nothing to factorize away.
341
+
342
+ Every panel is also available on its own — `plot_likelihood`, `plot_trial`,
343
+ `plot_regret`, `plot_relative_accuracy`, `plot_regret_vs_accuracy`,
344
+ `plot_map_agreement`, `plot_belief_profile`, `plot_confidence_density` — and each
345
+ takes `ax` (or `axes`, or a `subplot_spec`) so you can compose your own layout.
346
+
347
+ ### Watching an episode
348
+
349
+ `animate_episode` is the `render()` of this package. In a notebook, make it the
350
+ last expression in a cell and it plays — no ffmpeg, no `to_jshtml`, no backend to
351
+ configure:
352
+
353
+ ```python
354
+ from coggrid.viz import animate_episode
355
+
356
+ animate_episode(batch, traces) # plays inline
357
+ animate_episode(batch, traces).save("docs/images/episode") # or write a GIF
358
+ animate_episode(batch, traces).to_html() # or a scrub player
359
+ ```
360
+
361
+ Panels are laid out in rows — realization grids on top, anything that evolves
362
+ over time underneath:
363
+
364
+ | Row | Panels |
365
+ | --- | --- |
366
+ | `GRID_PANELS` | `joint`, `naive (factorized)`, and `joint − naive` — where factorizing moves probability mass. Each marks the truth in green and the observer's current mode with a white ring, and the axes name which variable is the **goal** (green) and which is **context** (orange). |
367
+ | `TRACE_PANELS` | the evidence stream, revealed up to the current step |
368
+
369
+ `animate_episode` takes `episode` (which episode to play, default 0), `fps`
370
+ (playback speed, default 6), and two arguments that decide what is drawn:
371
+
372
+ - **`extended=True`** — used for the animation at the top of this page — adds the
373
+ goal variable's marginal belief and two panels relating factorization regret to
374
+ dis-entanglement, one accumulating and one per step.
375
+ - **`panels`** — a list of rows, each row a list of panel functions, replacing the
376
+ default layout entirely. This is the extension point below.
377
+
378
+ **Adding your own panel.** A panel draws its furniture into an axis and returns
379
+ an updater called with the timestep, so extending the animation never means
380
+ editing the library:
381
+
382
+ ```python
383
+ from coggrid.viz import GRID_PANELS, TRACE_PANELS, animate_episode
384
+
385
+ def p_correct_panel(ax, view):
386
+ (line,) = ax.plot([], [])
387
+ p_correct = view.traces["joint"].p_correct[view.episode]
388
+ ax.set(xlim=(0, view.n_steps - 1), ylim=(0, 1), title="P(true value)")
389
+
390
+ def update(t):
391
+ line.set_data(range(t + 1), p_correct[: t + 1])
392
+
393
+ return update
394
+
395
+ animate_episode(batch, traces, panels=[GRID_PANELS, [*TRACE_PANELS, p_correct_panel]])
396
+ ```
397
+
398
+ `view` is an `EpisodeView` — the batch, traces, episode index and palette, plus
399
+ the derived arrays panels keep needing (`joint_grid`, `naive_grid`,
400
+ `observations`, `truth`, `goal_regret`, `disentanglement`).
401
+
402
+ ```bash
403
+ python examples/04_figures.py --out docs/images/ # the four static figures
404
+ python examples/05_animation.py --out docs/images/ # an episode played back, as a GIF
405
+ ```
406
+
407
+ ---
408
+
409
+ ## Customizing the generative model
410
+
411
+ Five stages are swappable. Pass a function to `World`; you never edit the
412
+ library.
413
+
414
+ ```python
415
+ def codebook_embeddings(cfg, rng):
416
+ """An EmbeddingSource: (cfg, rng) -> (keys, queries)."""
417
+ ...
418
+
419
+ world = World(cfg, embeddings=codebook_embeddings)
420
+ ```
421
+
422
+ | Argument | Signature | Replace it to change |
423
+ | --- | --- | --- |
424
+ | `embeddings` | `EmbeddingSource` | How latent variables relate to each other |
425
+ | `contexts` | `ContextSampler` | Which variables are active; the split structure |
426
+ | `likelihood` | `LikelihoodModel` | The *form* of the interaction (e.g. add a 3-way term) |
427
+ | `realizations` | `RealizationSampler` | The prior over latent values |
428
+ | `observations` | `ObservationModel` | Non-stationary, correlated or continuous observations |
429
+
430
+ Signatures are at the top of [`src/coggrid/world.py`](https://github.com/johnschwarcz/coggrid/blob/main/src/coggrid/world.py).
431
+ Because each returns its outputs rather than mutating shared state, a replacement
432
+ can be unit-tested on its own.
433
+
434
+ ---
435
+
436
+ ## Memory
437
+
438
+ The joint likelihood is `batch_size x n_observations x n_realizations ** n_contexts`.
439
+ That last term grows fast, so check before you run:
440
+
441
+ ```python
442
+ >>> print(CogGridConfig(n_contexts=4, n_realizations=20).memory_report(1000))
443
+ 1000 episodes x 5 channels x 20^4 realizations
444
+ joint likelihood : 6.0 GiB
445
+ joint belief : 35.8 GiB
446
+ peak (approx) : 41.7 GiB
447
+ ```
448
+
449
+ A `ResourceWarning` fires automatically above 2 GiB.
450
+
451
+ ---
452
+
453
+ ## Layout
454
+
455
+ ```
456
+ src/coggrid/
457
+ ├── config.py CogGridConfig — every tunable, validated, immutable
458
+ ├── generative.py the generative model as pure functions
459
+ ├── world.py World (config + embeddings), EpisodeBatch, and the five
460
+ │ swappable generative-stage signatures
461
+ ├── observers.py ideal-observer baselines, factorization regret,
462
+ │ dis-entanglement
463
+ ├── env.py CogGridEnv, the Gymnasium environment
464
+ └── viz/ plotting and animation
465
+ examples/ five runnable scripts
466
+ docs/make_assets.py regenerates the images in this README
467
+ tests/ behavioral tests, plus the docstring examples
468
+ ```
469
+
470
+ Numerical notes worth knowing:
471
+
472
+ - **Randomness** is an explicit `numpy.random.Generator` throughout — no global
473
+ `np.random` — so any batch is reproducible from a seed.
474
+ - **Inference** accumulates log-likelihoods with a cumulative sum rather than
475
+ multiplying and renormalizing per step, so it does not underflow at long
476
+ horizons.
477
+ - **Per-step belief updates** are recovered by differencing *log* beliefs.
478
+ Beliefs routinely fall below `1e-12` once an observer is confident, and
479
+ dividing two such numbers loses most of the significant digits.
480
+ - **Rates** are clipped away from 0 and 1 before taking logs, and the sigmoid is
481
+ overflow-safe.
482
+
483
+ ---
484
+
485
+ ## Related repositories
486
+
487
+ | Repository | Contents |
488
+ | --- | --- |
489
+ | **this one** (`coggrid`) | The environment and ideal-observer baselines. No networks, no torch. |
490
+ | [`CognitiveGridworld`](https://github.com/johnschwarcz/CognitiveGridworld) | The reference implementation from the paper: environment *and* trained networks, together, as published. |
491
+
492
+ Use **this** repo if you want the task. Use the paper repo to reproduce published
493
+ results.
494
+
495
+ ---
496
+
497
+ ## Citation
498
+
499
+ ```bibtex
500
+ @article{schwarcz2026factorization,
501
+ title = {Factorization Regret mediates compositional generalization in latent space},
502
+ author = {Schwarcz, John},
503
+ year = {2026},
504
+ eprint = {2603.27134},
505
+ archivePrefix = {arXiv}
506
+ }
507
+ ```
508
+
509
+ ## License
510
+
511
+ MIT — see [LICENSE](https://github.com/johnschwarcz/coggrid/blob/main/LICENSE).