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.
- coggrid-0.2.0/.gitignore +13 -0
- coggrid-0.2.0/LICENSE +21 -0
- coggrid-0.2.0/PKG-INFO +511 -0
- coggrid-0.2.0/README.md +477 -0
- coggrid-0.2.0/docs/make_assets.py +89 -0
- coggrid-0.2.0/examples/01_quickstart.py +34 -0
- coggrid-0.2.0/examples/02_gym_loop.py +75 -0
- coggrid-0.2.0/examples/03_customize.py +44 -0
- coggrid-0.2.0/examples/04_figures.py +32 -0
- coggrid-0.2.0/examples/05_animation.py +51 -0
- coggrid-0.2.0/pyproject.toml +76 -0
- coggrid-0.2.0/src/coggrid/__init__.py +81 -0
- coggrid-0.2.0/src/coggrid/config.py +240 -0
- coggrid-0.2.0/src/coggrid/env.py +287 -0
- coggrid-0.2.0/src/coggrid/generative.py +435 -0
- coggrid-0.2.0/src/coggrid/observers.py +347 -0
- coggrid-0.2.0/src/coggrid/viz/__init__.py +102 -0
- coggrid-0.2.0/src/coggrid/viz/animate.py +1000 -0
- coggrid-0.2.0/src/coggrid/viz/plots.py +1603 -0
- coggrid-0.2.0/src/coggrid/viz/style.py +80 -0
- coggrid-0.2.0/src/coggrid/world.py +378 -0
- coggrid-0.2.0/tests/test_environment.py +398 -0
- coggrid-0.2.0/tests/test_readme.py +236 -0
- coggrid-0.2.0/tests/test_scripts.py +124 -0
- coggrid-0.2.0/tests/test_viz.py +529 -0
coggrid-0.2.0/.gitignore
ADDED
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
|
+
[](https://github.com/johnschwarcz/coggrid/actions/workflows/ci.yml)
|
|
38
|
+
[](https://pypi.org/project/coggrid/)
|
|
39
|
+
[](https://www.python.org/downloads/)
|
|
40
|
+
[](https://github.com/johnschwarcz/coggrid/blob/main/LICENSE)
|
|
41
|
+
[](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).
|