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.
- exege/__init__.py +17 -0
- exege/_cli.py +89 -0
- exege/_render.py +31 -0
- exege/adapters/__init__.py +6 -0
- exege/adapters/bundle_dir.py +315 -0
- exege/adapters/latent_archive.py +561 -0
- exege/adapters/toy_dynamics.py +168 -0
- exege/app/__init__.py +14 -0
- exege/app/cli.py +68 -0
- exege/app/config.py +35 -0
- exege/app/latent.py +722 -0
- exege/app/main.py +19 -0
- exege/app/record.py +228 -0
- exege/app/theme.py +11 -0
- exege/core/__init__.py +21 -0
- exege/core/errors.py +25 -0
- exege/core/extras.py +77 -0
- exege/core/registry.py +138 -0
- exege/figures/__init__.py +28 -0
- exege/figures/bars.py +137 -0
- exege/figures/grids.py +149 -0
- exege/figures/maps.py +234 -0
- exege/figures/series.py +91 -0
- exege/latents/__init__.py +315 -0
- exege/latents/analysis.py +402 -0
- exege/latents/basis.py +521 -0
- exege/latents/cli.py +1137 -0
- exege/latents/evaluate.py +851 -0
- exege/latents/features.py +318 -0
- exege/latents/grid.py +190 -0
- exege/latents/record.py +324 -0
- exege/latents/samples.py +188 -0
- exege/latents/source.py +480 -0
- exege/latents/steering.py +530 -0
- exege/latents/through.py +770 -0
- exege/latents/toy.py +276 -0
- exege/nn/__init__.py +14 -0
- exege/nn/cli.py +81 -0
- exege/nn/sae.py +192 -0
- exege/nn/train.py +200 -0
- exege_core-0.6.0.dist-info/METADATA +120 -0
- exege_core-0.6.0.dist-info/RECORD +46 -0
- exege_core-0.6.0.dist-info/WHEEL +4 -0
- exege_core-0.6.0.dist-info/entry_points.txt +7 -0
- exege_core-0.6.0.dist-info/licenses/LICENSE +28 -0
- 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
|