exege-core 0.6.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.
Files changed (46) hide show
  1. exege/__init__.py +17 -0
  2. exege/_cli.py +89 -0
  3. exege/_render.py +31 -0
  4. exege/adapters/__init__.py +6 -0
  5. exege/adapters/bundle_dir.py +315 -0
  6. exege/adapters/latent_archive.py +561 -0
  7. exege/adapters/toy_dynamics.py +168 -0
  8. exege/app/__init__.py +14 -0
  9. exege/app/cli.py +68 -0
  10. exege/app/config.py +35 -0
  11. exege/app/latent.py +722 -0
  12. exege/app/main.py +19 -0
  13. exege/app/record.py +228 -0
  14. exege/app/theme.py +11 -0
  15. exege/core/__init__.py +21 -0
  16. exege/core/errors.py +25 -0
  17. exege/core/extras.py +77 -0
  18. exege/core/registry.py +138 -0
  19. exege/figures/__init__.py +28 -0
  20. exege/figures/bars.py +137 -0
  21. exege/figures/grids.py +149 -0
  22. exege/figures/maps.py +234 -0
  23. exege/figures/series.py +91 -0
  24. exege/latents/__init__.py +315 -0
  25. exege/latents/analysis.py +402 -0
  26. exege/latents/basis.py +521 -0
  27. exege/latents/cli.py +1137 -0
  28. exege/latents/evaluate.py +851 -0
  29. exege/latents/features.py +318 -0
  30. exege/latents/grid.py +190 -0
  31. exege/latents/record.py +324 -0
  32. exege/latents/samples.py +188 -0
  33. exege/latents/source.py +480 -0
  34. exege/latents/steering.py +530 -0
  35. exege/latents/through.py +770 -0
  36. exege/latents/toy.py +276 -0
  37. exege/nn/__init__.py +14 -0
  38. exege/nn/cli.py +81 -0
  39. exege/nn/sae.py +192 -0
  40. exege/nn/train.py +200 -0
  41. exege_core-0.6.0.dist-info/METADATA +120 -0
  42. exege_core-0.6.0.dist-info/RECORD +46 -0
  43. exege_core-0.6.0.dist-info/WHEEL +4 -0
  44. exege_core-0.6.0.dist-info/entry_points.txt +7 -0
  45. exege_core-0.6.0.dist-info/licenses/LICENSE +28 -0
  46. exege_core-0.6.0.dist-info/licenses/NOTICE +40 -0
exege/__init__.py ADDED
@@ -0,0 +1,17 @@
1
+ """exege -- tools for understanding and evaluating scientific machine-learning models.
2
+
3
+ Domains, and the presentation downstream of them:
4
+
5
+ - ``exege.latents`` what emulators hold inside: the latent space, on a grid
6
+ - ``exege.nn`` torch modules trained on those latents: a sparse autoencoder
7
+ - ``exege.figures`` figures of what ``latents`` computes, with no web framework
8
+ - ``exege.app`` a local web app over all of the above; nothing imports it
9
+
10
+ Everything framework-specific lives in ``exege.adapters``. See ``AGENTS.md``.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ __version__ = "0.6.0"
16
+
17
+ __all__ = ["__version__"]
exege/_cli.py ADDED
@@ -0,0 +1,89 @@
1
+ """Top-level ``exege`` command.
2
+
3
+ Subcommands load lazily, so ``exege --help`` and one command never pay for (or
4
+ require) what stands behind another.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import logging
10
+ import sys
11
+ from importlib import import_module
12
+ from importlib.metadata import entry_points
13
+
14
+ import click
15
+
16
+ from exege import __version__
17
+ from exege.core.errors import ExegeError
18
+
19
+ log = logging.getLogger(__name__)
20
+
21
+ # Shipped commands, by import path. Listing them for --help imports every cli
22
+ # module, so each must stay importable on the base tier: anything heavier is
23
+ # imported inside the command that needs it (tests/test_purity.py holds them to it).
24
+ _COMMANDS = {
25
+ "latents": "exege.latents.cli:latents",
26
+ "nn": "exege.nn.cli:nn",
27
+ "app": "exege.app.cli:app",
28
+ }
29
+
30
+ # A separate distribution adds a command by registering a ``click.Command``
31
+ # under this entry-point group. Shipped names win a clash.
32
+ _GROUP = "exege.commands"
33
+
34
+
35
+ class _LazyGroup(click.Group):
36
+ def list_commands(self, ctx: click.Context) -> list[str]:
37
+ return sorted(set(_COMMANDS) | {ep.name for ep in entry_points(group=_GROUP)})
38
+
39
+ def get_command(self, ctx: click.Context, name: str) -> click.Command | None:
40
+ if name in _COMMANDS:
41
+ module, _, attr = _COMMANDS[name].partition(":")
42
+ return getattr(import_module(module), attr)
43
+ for ep in entry_points(group=_GROUP):
44
+ if ep.name == name:
45
+ try:
46
+ return ep.load()
47
+ except Exception as exc: # one broken plugin must not take --help down
48
+ log.warning("could not load command %r from %s: %s", name, ep.value, exc)
49
+ return None
50
+
51
+ def invoke(self, ctx: click.Context) -> object:
52
+ """Deliberate errors read as one clear line, not a traceback -- here, so
53
+ that a test runner or an embedding program sees what a terminal does.
54
+ ``--debug`` keeps the traceback."""
55
+ try:
56
+ return super().invoke(ctx)
57
+ except ExegeError as exc:
58
+ if ctx.params.get("debug"):
59
+ raise
60
+ raise click.ClickException(str(exc)) from exc
61
+
62
+
63
+ @click.group(cls=_LazyGroup)
64
+ @click.version_option(__version__, prog_name="exege")
65
+ @click.option("--debug", is_flag=True, help="Verbose logging.")
66
+ def cli(debug: bool) -> None:
67
+ """Tools for understanding and evaluating scientific machine-learning models."""
68
+ logging.basicConfig(
69
+ level=logging.DEBUG if debug else logging.WARNING,
70
+ format="%(levelname)s %(name)s: %(message)s",
71
+ )
72
+
73
+
74
+ def main() -> None:
75
+ try:
76
+ # Not standalone, so that an interrupt exits 130 without click's "Aborted!".
77
+ # click then *returns* the code of a `ctx.exit(n)` instead of exiting with it.
78
+ code = cli.main(standalone_mode=False)
79
+ except click.ClickException as exc:
80
+ exc.show()
81
+ sys.exit(exc.exit_code)
82
+ except click.Abort:
83
+ sys.exit(130)
84
+ if isinstance(code, int) and code:
85
+ sys.exit(code)
86
+
87
+
88
+ if __name__ == "__main__":
89
+ main()
exege/_render.py ADDED
@@ -0,0 +1,31 @@
1
+ """Plain-text rendering. Presentation only, and stdlib only."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Iterable, Mapping, Sequence
6
+ from typing import Any
7
+
8
+
9
+ def _cell(value: Any) -> str:
10
+ return "-" if value is None or value == "" else str(value)
11
+
12
+
13
+ def table(rows: Sequence[Mapping[str, Any]], columns: Iterable[str] | None = None) -> str:
14
+ """Left-aligned fixed-width table. Returns '' for no rows so callers can
15
+ distinguish 'nothing found' from 'a header with nothing under it'."""
16
+ if not rows:
17
+ return ""
18
+ cols = list(columns) if columns is not None else list(rows[0].keys())
19
+ widths = {c: max(len(c), *(len(_cell(r.get(c))) for r in rows)) for c in cols}
20
+ out = [" ".join(c.upper().ljust(widths[c]) for c in cols).rstrip()]
21
+ for row in rows:
22
+ out.append(" ".join(_cell(row.get(c)).ljust(widths[c]) for c in cols).rstrip())
23
+ return "\n".join(out)
24
+
25
+
26
+ def pairs(mapping: Mapping[str, Any]) -> str:
27
+ """Aligned ``key: value`` block."""
28
+ if not mapping:
29
+ return ""
30
+ width = max(len(k) for k in mapping)
31
+ return "\n".join(f"{k.ljust(width)} {_cell(v)}" for k, v in mapping.items())
@@ -0,0 +1,6 @@
1
+ """Adapters: every assumption about a specific framework lives here.
2
+
3
+ Adapters are resolved through the ``exege.adapters`` entry-point group, so a new
4
+ system ships as a new module (or a separate distribution) and core never changes.
5
+ Heavy dependencies belong here, behind an extra -- never in the base tier.
6
+ """
@@ -0,0 +1,315 @@
1
+ """Read a bundle: a second latent layout, unlike the archive on purpose.
2
+
3
+ It exists to show that a system with another file layout, another vocabulary and another
4
+ dimension order is a new module here and an entry point, and nothing else. Where an archive
5
+ keeps ``(n_times, n_nodes, n_channels)`` in one file per layer, a bundle keeps channels
6
+ first, on a grid that is two-dimensional, one file per level *and* time::
7
+
8
+ bundle.json what the system is, its clock, its levels, its grid, notes
9
+ coords.npz lat (n_lat,), lon (n_lon,); optional ocean (n_lat, n_lon) bool,
10
+ True where the model says something; optional cell_area
11
+ levels/L00/T0000.npy (width, n_lat, n_lon), any float dtype
12
+ fields/<name>.npy optional (n_stamps, n_lat, n_lon): physical fields
13
+
14
+ Nodes are the grid read in C order, ``node = i_lat * n_lon + i_lon``, which is what every
15
+ ``Grid`` of exege means. Files are memory-mapped, and a read takes the channels asked for
16
+ first and the nodes second, so a region of one level touches only the pages it needs.
17
+
18
+ ``bundle.json`` says::
19
+
20
+ {"format": "exege-bundle", "version": 1,
21
+ "system": {"name": ..., "part": ..., "weights": ...},
22
+ "clock": {"calendar": "noleap", "step_seconds": 21600, "stamps": [...],
23
+ "field_stamps": [...]},
24
+ "levels": [{"id": 0, "name": ..., "width": 16, "network_position": 4}],
25
+ "grid": [n_lat, n_lon], "fields": [...], "notes": {...}}
26
+
27
+ ``write`` is the other half, with the signature ``write_archive`` has, so the toy emulator
28
+ (``write_toy(path, adapter="bundle-dir")``) and anything else that writes archives writes
29
+ bundles too, and the reader is tested against it.
30
+ """
31
+
32
+ from __future__ import annotations
33
+
34
+ import json
35
+ import shutil
36
+ from collections.abc import Mapping, Sequence
37
+ from pathlib import Path
38
+ from typing import Any
39
+
40
+ from exege.core.errors import AdapterError, RequestError
41
+ from exege.core.extras import missing_extra
42
+ from exege.latents.grid import Grid
43
+ from exege.latents.source import LatentInfo, LayerInfo, selection
44
+
45
+ try:
46
+ import numpy as np
47
+ except ImportError as exc:
48
+ raise missing_extra("numpy", "latents") from exc
49
+
50
+ MANIFEST = "bundle.json"
51
+ COORDS = "coords.npz"
52
+ FORMAT = "exege-bundle"
53
+ # What bundles were tagged before the package was renamed from xaig; still read.
54
+ OLD_FORMATS = ("xaig-bundle",)
55
+ VERSION = 1
56
+
57
+
58
+ def _level_file(level: int, stamp: int) -> str:
59
+ return f"levels/L{level:02d}/T{stamp:04d}.npy"
60
+
61
+
62
+ class BundleDir:
63
+ """A ``LatentSource``, and ``ReferenceFields``, over one bundle directory.
64
+
65
+ ``unmasked`` ignores the bundle's ocean mask and treats every node as valid, for a
66
+ reader who wants what the model wrote over land too."""
67
+
68
+ def __init__(self, path: str | Path, *, unmasked: bool = False) -> None:
69
+ self.path = Path(path)
70
+ self.unmasked = bool(unmasked)
71
+ if not self.path.is_dir():
72
+ raise AdapterError(f"no such directory: {self.path}")
73
+ manifest_path = self.path / MANIFEST
74
+ if not manifest_path.is_file():
75
+ raise AdapterError(f"not a bundle (no {MANIFEST}): {self.path}")
76
+ try:
77
+ manifest = json.loads(manifest_path.read_text())
78
+ if manifest["format"] not in (FORMAT, *OLD_FORMATS) or manifest["version"] != VERSION:
79
+ raise AdapterError(
80
+ f"{manifest_path}: format {manifest['format']!r} version "
81
+ f"{manifest['version']!r}; this reads {FORMAT!r} version {VERSION}"
82
+ )
83
+ self._shape = (int(manifest["grid"][0]), int(manifest["grid"][1]))
84
+ clock, system = manifest["clock"], manifest.get("system") or {}
85
+ stamps = tuple(str(t) for t in clock["stamps"])
86
+ levels = [
87
+ LayerInfo(
88
+ int(level["id"]),
89
+ str(level.get("name", "")),
90
+ int(level["width"]),
91
+ None
92
+ if level.get("network_position") is None
93
+ else int(level["network_position"]),
94
+ )
95
+ for level in manifest["levels"]
96
+ ]
97
+ except (json.JSONDecodeError, KeyError, IndexError, TypeError, ValueError) as exc:
98
+ raise AdapterError(f"{manifest_path}: unreadable manifest ({exc!r})") from exc
99
+ if not levels or not stamps:
100
+ raise AdapterError(f"{manifest_path}: a bundle needs levels and stamps")
101
+ if len({level.index for level in levels}) != len(levels):
102
+ raise AdapterError(f"{manifest_path}: a level id appears more than once")
103
+ notes = manifest.get("notes") or {}
104
+ if not isinstance(notes, dict):
105
+ raise AdapterError(f"{manifest_path}: 'notes' must be a mapping")
106
+ self._field_names = tuple(str(n) for n in manifest.get("fields") or ())
107
+ self._field_stamps = tuple(str(t) for t in clock.get("field_stamps") or stamps)
108
+ self._grid: Grid | None = None
109
+ step = clock.get("step_seconds")
110
+ self._info = LatentInfo(
111
+ source=str(self.path),
112
+ times=stamps,
113
+ layers=tuple(levels),
114
+ n_nodes=self._shape[0] * self._shape[1],
115
+ model=system.get("name"),
116
+ component=system.get("part"),
117
+ checkpoint=system.get("weights"),
118
+ calendar=clock.get("calendar"),
119
+ timestep_seconds=None if step is None else int(step),
120
+ experiment=notes,
121
+ options={"unmasked": True} if self.unmasked else {},
122
+ )
123
+
124
+ @staticmethod
125
+ def write(path: str | Path, **contents: Any) -> Path:
126
+ """``write_bundle``, reachable from the class the registry hands out."""
127
+ return write_bundle(path, **contents)
128
+
129
+ def info(self) -> LatentInfo:
130
+ return self._info
131
+
132
+ def grid(self) -> Grid:
133
+ if self._grid is None:
134
+ self._grid = self._read_grid()
135
+ return self._grid
136
+
137
+ def _read_grid(self) -> Grid:
138
+ file = self.path / COORDS
139
+ if not file.is_file():
140
+ raise AdapterError(f"{self.path}: no {COORDS}")
141
+ with np.load(file) as stored:
142
+ names = set(stored.files)
143
+ if not {"lat", "lon"} <= names:
144
+ raise AdapterError(f"{file}: needs 'lat' and 'lon', has {sorted(names)}")
145
+ lat_1d, lon_1d = (np.asarray(stored[k], dtype=np.float64) for k in ("lat", "lon"))
146
+ ocean = np.asarray(stored["ocean"], dtype=bool) if "ocean" in names else None
147
+ area = (
148
+ np.asarray(stored["cell_area"], dtype=np.float64) if "cell_area" in names else None
149
+ )
150
+ if (lat_1d.size, lon_1d.size) != self._shape:
151
+ raise AdapterError(
152
+ f"{file}: {lat_1d.size} latitudes and {lon_1d.size} longitudes, but the "
153
+ f"manifest says a grid of {self._shape}"
154
+ )
155
+ lat, lon = np.meshgrid(lat_1d, lon_1d, indexing="ij")
156
+ mask = None if self.unmasked or ocean is None else ocean.ravel()
157
+ try:
158
+ return Grid(
159
+ lat=lat.ravel(), lon=lon.ravel(), shape=self._shape, mask=mask,
160
+ area=None if area is None else area.ravel(),
161
+ ) # fmt: skip
162
+ except ValueError as exc:
163
+ raise AdapterError(f"{file}: {exc}") from exc
164
+
165
+ # -- ReferenceFields ------------------------------------------------------
166
+
167
+ def field_names(self) -> tuple[str, ...]:
168
+ return self._field_names
169
+
170
+ def field(self, name: str, time: str | int, lead: int = 0) -> np.ndarray:
171
+ label = self._info.times[self._info.time_index(time)]
172
+ if name not in self._field_names:
173
+ known = ", ".join(self._field_names) or "none"
174
+ raise RequestError(f"no field {name!r} in {self.path}; fields are {known}")
175
+ if label not in self._field_stamps:
176
+ raise RequestError(f"{self.path}: the fields hold no time {label!r}")
177
+ index = self._field_stamps.index(label) + lead
178
+ if not 0 <= index < len(self._field_stamps):
179
+ raise RequestError(f"{self.path}: the fields end before {lead:+d} time(s) from {label}")
180
+ stored = self._open(f"fields/{name}.npy", (len(self._field_stamps), *self._shape))
181
+ return np.asarray(stored[index], dtype=np.float64).ravel()
182
+
183
+ # -- LatentSource ---------------------------------------------------------
184
+
185
+ def _open(self, name: str, expected: tuple[int, ...]) -> np.ndarray:
186
+ file = self.path / name
187
+ if not file.is_file():
188
+ raise AdapterError(f"{file} is missing")
189
+ array = np.load(file, mmap_mode="r")
190
+ if array.shape != expected:
191
+ raise AdapterError(f"{file}: shape {array.shape}, but the manifest implies {expected}")
192
+ return array
193
+
194
+ def load(
195
+ self,
196
+ time: str | int,
197
+ layer: int,
198
+ channels: Sequence[int] | None = None,
199
+ nodes: Sequence[int] | None = None,
200
+ ) -> np.ndarray:
201
+ width = self._info.layer(layer).n_channels
202
+ stamp = self._info.time_index(time)
203
+ block = self._open(_level_file(layer, stamp), (width, *self._shape))
204
+ block = block.reshape(width, self._info.n_nodes) # a view: the file is C-ordered
205
+ # Channels (the leading axis) first, then nodes: on a memory map the rows asked
206
+ # for are the only pages read.
207
+ held = f"level {layer} holds {width} channels x {self._info.n_nodes} nodes"
208
+ if channels is not None:
209
+ block = block[selection(channels, width, "channel", held)]
210
+ if nodes is not None:
211
+ block = block[:, selection(nodes, block.shape[1], "node", held)]
212
+ return np.array(block.T, dtype=np.float32, order="C") # a copy, nodes first
213
+
214
+
215
+ def write_bundle(
216
+ path: str | Path,
217
+ *,
218
+ grid: Grid,
219
+ times: Sequence[str],
220
+ layers: Sequence[tuple[str, np.ndarray]],
221
+ network_layers: Sequence[int] | None = None,
222
+ fields: Mapping[str, np.ndarray] | None = None,
223
+ field_times: Sequence[str] | None = None,
224
+ model: str | None = None,
225
+ component: str | None = None,
226
+ checkpoint: str | None = None,
227
+ calendar: str | None = None,
228
+ timestep_seconds: int | None = None,
229
+ experiment: Mapping[str, Any] | None = None,
230
+ overwrite: bool = False,
231
+ ) -> Path:
232
+ """Write one bundle directory, in float32, from what ``write_archive`` is given:
233
+ ``layers`` as ``(label, array)`` pairs of ``(n_times, n_nodes, n_channels)``,
234
+ ``fields`` as ``(n_field_times, n_nodes)``. Only a grid of a shape (latitudes by
235
+ longitudes) fits the layout. A demonstration, so written in place, not staged."""
236
+ out = Path(path)
237
+ times = [str(t) for t in times]
238
+ if grid.shape is None:
239
+ raise RequestError("a bundle holds a latitude-longitude grid; this one is a mesh")
240
+ n_lat, n_lon = grid.shape
241
+ if not times or len(set(times)) != len(times):
242
+ raise RequestError("a bundle needs at least one time, and each only once")
243
+ if not layers:
244
+ raise RequestError("a bundle needs at least one level")
245
+ if network_layers is not None and len(network_layers) != len(layers):
246
+ raise RequestError(f"{len(network_layers)} network layer(s) for {len(layers)} level(s)")
247
+ for index, (label, array) in enumerate(layers):
248
+ if array.ndim != 3 or array.shape[:2] != (len(times), grid.n_nodes):
249
+ raise RequestError(
250
+ f"level {index} ({label!r}) has shape {array.shape}; expected "
251
+ f"({len(times)} times, {grid.n_nodes} nodes, channels)"
252
+ )
253
+ if array.dtype.kind not in "biuf":
254
+ raise RequestError(f"level {index} must hold real numeric values")
255
+ fields = dict(fields or {})
256
+ stamps = [str(t) for t in (field_times if field_times is not None else times)]
257
+ if fields:
258
+ if len(set(stamps)) != len(stamps) or not set(times) <= set(stamps):
259
+ raise RequestError("field times must be unique and hold every latent time")
260
+ for name, values in fields.items():
261
+ if values.shape != (len(stamps), grid.n_nodes):
262
+ raise RequestError(
263
+ f"field {name!r} has shape {values.shape}; expected "
264
+ f"({len(stamps)} times, {grid.n_nodes} nodes)"
265
+ )
266
+ if out.is_symlink() or (out.exists() and not out.is_dir()):
267
+ raise RequestError(f"{out} must be a directory, not a file or symlink")
268
+ if out.exists() and any(out.iterdir()) and not overwrite:
269
+ raise RequestError(f"{out} is not empty; pass overwrite=True to replace it")
270
+
271
+ levels = []
272
+ for index, (label, array) in enumerate(layers):
273
+ level = {"id": index, "name": str(label), "width": int(array.shape[2])}
274
+ if network_layers is not None:
275
+ level["network_position"] = int(network_layers[index])
276
+ levels.append(level)
277
+ clock = {"calendar": calendar, "step_seconds": timestep_seconds, "stamps": times}
278
+ if fields:
279
+ clock["field_stamps"] = stamps
280
+ manifest: dict[str, Any] = {
281
+ "format": FORMAT, "version": VERSION,
282
+ "system": {"name": model, "part": component, "weights": checkpoint},
283
+ "clock": {k: v for k, v in clock.items() if v is not None},
284
+ "levels": levels, "grid": [n_lat, n_lon], "fields": sorted(fields),
285
+ "notes": dict(experiment or {}),
286
+ } # fmt: skip
287
+ try:
288
+ metadata = json.dumps(manifest, indent=2)
289
+ except (TypeError, ValueError) as exc:
290
+ raise RequestError(f"bundle metadata must be JSON serializable: {exc}") from exc
291
+
292
+ coords: dict[str, np.ndarray] = {
293
+ "lat": grid.lat.reshape(grid.shape)[:, 0], "lon": grid.lon.reshape(grid.shape)[0, :]
294
+ } # fmt: skip
295
+ if grid.mask is not None:
296
+ coords["ocean"] = np.asarray(grid.mask, dtype=bool).reshape(grid.shape)
297
+ if grid.area is not None:
298
+ coords["cell_area"] = np.asarray(grid.area).reshape(grid.shape)
299
+
300
+ if out.exists():
301
+ shutil.rmtree(out)
302
+ (out / "levels").mkdir(parents=True)
303
+ np.savez(out / COORDS, **coords)
304
+ for index, (_, array) in enumerate(layers):
305
+ (out / "levels" / f"L{index:02d}").mkdir()
306
+ for stamp in range(len(times)):
307
+ block = array[stamp].T.reshape(array.shape[2], n_lat, n_lon)
308
+ np.save(out / _level_file(index, stamp), block.astype(np.float32))
309
+ if fields:
310
+ (out / "fields").mkdir()
311
+ for name, values in fields.items():
312
+ block = values.reshape(len(stamps), n_lat, n_lon).astype(np.float32)
313
+ np.save(out / "fields" / f"{name}.npy", block)
314
+ (out / MANIFEST).write_text(metadata)
315
+ return out