coordax 0.1.1__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.
- coordax/__init__.py +47 -0
- coordax/coordinate_systems.py +569 -0
- coordax/coordinate_systems_test.py +299 -0
- coordax/fields.py +718 -0
- coordax/fields_test.py +668 -0
- coordax/jax_datetime_integration_test.py +64 -0
- coordax/named_axes.py +938 -0
- coordax/named_axes_test.py +963 -0
- coordax/ndarrays.py +147 -0
- coordax/testing.py +77 -0
- coordax/utils.py +23 -0
- coordax/xarray_test.py +234 -0
- coordax-0.1.1.dist-info/METADATA +234 -0
- coordax-0.1.1.dist-info/RECORD +17 -0
- coordax-0.1.1.dist-info/WHEEL +5 -0
- coordax-0.1.1.dist-info/licenses/LICENSE +202 -0
- coordax-0.1.1.dist-info/top_level.txt +1 -0
coordax/__init__.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
# Copyright 2024 Google LLC
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# https://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
# Note: import <name> as <name> is required for names to be exported.
|
|
16
|
+
# See PEP 484 & https://github.com/jax-ml/jax/issues/7570
|
|
17
|
+
# pylint: disable=g-multiple-import,useless-import-alias,g-importing-member
|
|
18
|
+
from coordax.coordinate_systems import (
|
|
19
|
+
CartesianProduct as CartesianProduct,
|
|
20
|
+
Coordinate as Coordinate,
|
|
21
|
+
DummyAxis as DummyAxis,
|
|
22
|
+
LabeledAxis as LabeledAxis,
|
|
23
|
+
Scalar as Scalar,
|
|
24
|
+
SizedAxis as SizedAxis,
|
|
25
|
+
NoCoordinateMatch as NoCoordinateMatch,
|
|
26
|
+
SelectedAxis as SelectedAxis,
|
|
27
|
+
canonicalize as canonicalize_coordinates,
|
|
28
|
+
compose as compose_coordinates,
|
|
29
|
+
from_xarray as coordinates_from_xarray,
|
|
30
|
+
)
|
|
31
|
+
from coordax.fields import (
|
|
32
|
+
Field as Field,
|
|
33
|
+
is_field as is_field,
|
|
34
|
+
tmp_axis_name as tmp_axis_name,
|
|
35
|
+
cmap as cmap,
|
|
36
|
+
get_coordinate as get_coordinate,
|
|
37
|
+
wrap_like as wrap_like,
|
|
38
|
+
wrap as wrap,
|
|
39
|
+
tag as tag,
|
|
40
|
+
untag as untag,
|
|
41
|
+
)
|
|
42
|
+
from coordax.ndarrays import (
|
|
43
|
+
NDArray as NDArray,
|
|
44
|
+
register_ndarray as register_ndarray,
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
__version__ = "0.1.1" # keep sync with pyproject.toml
|
|
@@ -0,0 +1,569 @@
|
|
|
1
|
+
# Copyright 2024 Google LLC
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# https://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
"""Coordinate systems for use on coordax.Field objects.
|
|
15
|
+
|
|
16
|
+
``Coordinate`` objects define a discretization schema, dimension names and
|
|
17
|
+
provide methods & coordinate field values to facilitate computations.
|
|
18
|
+
"""
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import abc
|
|
22
|
+
import collections
|
|
23
|
+
from collections.abc import Iterable
|
|
24
|
+
import dataclasses
|
|
25
|
+
import itertools
|
|
26
|
+
import typing
|
|
27
|
+
from typing import Any, Self, TYPE_CHECKING, TypeAlias, TypeVar
|
|
28
|
+
|
|
29
|
+
from coordax import utils
|
|
30
|
+
import jax
|
|
31
|
+
import numpy as np
|
|
32
|
+
# TODO(shoyer): consider making Xarray an optional dependency of core Coordax
|
|
33
|
+
import xarray
|
|
34
|
+
|
|
35
|
+
if TYPE_CHECKING:
|
|
36
|
+
# import only under TYPE_CHECKING to avoid circular dependency
|
|
37
|
+
# pylint: disable=g-bad-import-order
|
|
38
|
+
from coordax import fields
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
Pytree: TypeAlias = Any
|
|
42
|
+
Sequence = collections.abc.Sequence
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@utils.export
|
|
46
|
+
@dataclasses.dataclass(frozen=True)
|
|
47
|
+
class NoCoordinateMatch:
|
|
48
|
+
"""For use when a coordinate does not match an xarray.Coordinates object."""
|
|
49
|
+
|
|
50
|
+
reason: str
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@utils.export
|
|
54
|
+
class Coordinate(abc.ABC):
|
|
55
|
+
"""Abstract class for coordinate objects.
|
|
56
|
+
|
|
57
|
+
Coordinate subclasses are expected to obey several invariants:
|
|
58
|
+
1. Dimension names may not be repeated: `len(set(dims)) == len(dims)`
|
|
59
|
+
2. All dimensions must be named: `len(shape) == len(dims)`
|
|
60
|
+
|
|
61
|
+
Every non-abstract Coordinate subclass must be registered as a "static"
|
|
62
|
+
pytree node, e.g., by decorating the class with
|
|
63
|
+
`@jax.tree_util.register_static`. Static pytrees nodes must implement
|
|
64
|
+
`__hash__` and `__eq__` according to the requirements of keys in Python
|
|
65
|
+
dictionaries. This is easiest to acheive with frozen dataclasses, but care
|
|
66
|
+
must be taken when working with np.ndarray attributes.
|
|
67
|
+
|
|
68
|
+
TODO(shoyer): add documentation examples, including a version using ArrayKey
|
|
69
|
+
to wrap np.ndarray attributions.
|
|
70
|
+
"""
|
|
71
|
+
|
|
72
|
+
@property
|
|
73
|
+
@abc.abstractmethod
|
|
74
|
+
def dims(self) -> tuple[str | None, ...]:
|
|
75
|
+
"""Dimension names of the coordinate.
|
|
76
|
+
|
|
77
|
+
All subclasses must return a tuple of dimension names as strings, with the
|
|
78
|
+
exception of `DummyAxis`.
|
|
79
|
+
"""
|
|
80
|
+
raise NotImplementedError()
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
@abc.abstractmethod
|
|
84
|
+
def shape(self) -> tuple[int, ...]:
|
|
85
|
+
"""Shape of the coordinate."""
|
|
86
|
+
raise NotImplementedError()
|
|
87
|
+
|
|
88
|
+
@property
|
|
89
|
+
@abc.abstractmethod
|
|
90
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
91
|
+
"""A maps from field names to their values."""
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def sizes(self) -> dict[str, int]:
|
|
95
|
+
"""Sizes of all dimensions on this coordinate."""
|
|
96
|
+
return {
|
|
97
|
+
dim: size for dim, size in zip(self.dims, self.shape) if dim is not None
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
@property
|
|
101
|
+
def ndim(self) -> int:
|
|
102
|
+
"""Dimensionality of the coordinate."""
|
|
103
|
+
return len(self.dims)
|
|
104
|
+
|
|
105
|
+
@property
|
|
106
|
+
def axes(self) -> tuple[Coordinate, ...]:
|
|
107
|
+
"""Tuple of one-dimensional Coordinate objects for each dimension."""
|
|
108
|
+
if self.ndim == 1:
|
|
109
|
+
return (self,)
|
|
110
|
+
else:
|
|
111
|
+
return tuple(SelectedAxis(self, i) for i in range(self.ndim))
|
|
112
|
+
|
|
113
|
+
def to_xarray(self) -> dict[str, xarray.Variable]:
|
|
114
|
+
"""Convert this coordinate into xarray variables."""
|
|
115
|
+
variables = {}
|
|
116
|
+
dims_set = {dim for dim in self.dims if dim is not None}
|
|
117
|
+
for name, coord_field in self.fields.items():
|
|
118
|
+
if set(coord_field.dims) <= dims_set:
|
|
119
|
+
# xarray.DataArray coordinate dimensions must be a subset of the
|
|
120
|
+
# dimensions of the associated DataArray, which is not necessarily a
|
|
121
|
+
# constraint for coordax.Field.
|
|
122
|
+
variables[name] = xarray.Variable(coord_field.dims, coord_field.data)
|
|
123
|
+
return variables
|
|
124
|
+
|
|
125
|
+
@classmethod
|
|
126
|
+
def from_xarray(
|
|
127
|
+
cls, dims: tuple[str, ...], coords: xarray.Coordinates
|
|
128
|
+
) -> Self | NoCoordinateMatch:
|
|
129
|
+
"""Construct a matching Coordax coordinate from xarray, if possible.
|
|
130
|
+
|
|
131
|
+
Args:
|
|
132
|
+
dims: tuple of dimension names. Only the leading dimensions should be
|
|
133
|
+
checks for a match.
|
|
134
|
+
coords: xarray.Coordinates object providing dimension sizes and coordinate
|
|
135
|
+
values.
|
|
136
|
+
|
|
137
|
+
Returns:
|
|
138
|
+
A matching instance of this coordinate or `NoCoordinateMatch` if this
|
|
139
|
+
coordinate does not match the xarray dimensions and coordinates.
|
|
140
|
+
"""
|
|
141
|
+
raise NotImplementedError('from_xarray not implemented')
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
@dataclasses.dataclass(frozen=True)
|
|
145
|
+
class ArrayKey:
|
|
146
|
+
"""Wrapper for a numpy array to make it hashable."""
|
|
147
|
+
|
|
148
|
+
value: np.ndarray
|
|
149
|
+
|
|
150
|
+
def __eq__(self, other):
|
|
151
|
+
return (
|
|
152
|
+
isinstance(self, ArrayKey)
|
|
153
|
+
and self.value.dtype == other.value.dtype
|
|
154
|
+
and self.value.shape == other.value.shape
|
|
155
|
+
and (self.value == other.value).all()
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
def __hash__(self) -> int:
|
|
159
|
+
return hash((self.value.shape, self.value.tobytes()))
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
@utils.export
|
|
163
|
+
@jax.tree_util.register_static
|
|
164
|
+
@dataclasses.dataclass(frozen=True)
|
|
165
|
+
class Scalar(Coordinate):
|
|
166
|
+
"""Zero dimensional sentinel coordinate used to label stand alone scalars."""
|
|
167
|
+
|
|
168
|
+
@property
|
|
169
|
+
def dims(self) -> tuple[str, ...]:
|
|
170
|
+
return ()
|
|
171
|
+
|
|
172
|
+
@property
|
|
173
|
+
def shape(self) -> tuple[int, ...]:
|
|
174
|
+
return ()
|
|
175
|
+
|
|
176
|
+
@property
|
|
177
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
178
|
+
return {}
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
@utils.export
|
|
182
|
+
@jax.tree_util.register_static
|
|
183
|
+
@dataclasses.dataclass(frozen=True)
|
|
184
|
+
class SelectedAxis(Coordinate):
|
|
185
|
+
"""Coordinate that exposes one dimension of a multidimensional coordinate."""
|
|
186
|
+
|
|
187
|
+
coordinate: Coordinate
|
|
188
|
+
axis: int
|
|
189
|
+
|
|
190
|
+
def __post_init__(self):
|
|
191
|
+
if self.axis >= self.coordinate.ndim:
|
|
192
|
+
raise ValueError(
|
|
193
|
+
f'Dimension {self.axis=} of {self.coordinate=} is out of bounds'
|
|
194
|
+
)
|
|
195
|
+
if self.coordinate.dims[self.axis] is None:
|
|
196
|
+
raise ValueError(
|
|
197
|
+
f'dimension {self.axis=} of {self.coordinate=} is not named'
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
@property
|
|
201
|
+
def dims(self) -> tuple[str, ...]:
|
|
202
|
+
"""Dimension names of the coordinate."""
|
|
203
|
+
dim = self.coordinate.dims[self.axis]
|
|
204
|
+
assert dim is not None
|
|
205
|
+
return (dim,)
|
|
206
|
+
|
|
207
|
+
@property
|
|
208
|
+
def shape(self) -> tuple[int, ...]:
|
|
209
|
+
"""Shape of the coordinate."""
|
|
210
|
+
return (self.coordinate.shape[self.axis],)
|
|
211
|
+
|
|
212
|
+
@property
|
|
213
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
214
|
+
"""A maps from field names to their values."""
|
|
215
|
+
return self.coordinate.fields
|
|
216
|
+
|
|
217
|
+
def __repr__(self):
|
|
218
|
+
return f'coordax.SelectedAxis({self.coordinate!r}, axis={self.axis})'
|
|
219
|
+
|
|
220
|
+
def to_xarray(self) -> dict[str, xarray.Variable]:
|
|
221
|
+
"""Convert this coordinate into xarray variables."""
|
|
222
|
+
# Override the default method to avoid restricting variables to only those
|
|
223
|
+
# along the selected axis.
|
|
224
|
+
return self.coordinate.to_xarray()
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def _expand_coordinates(*coordinates: Coordinate) -> tuple[Coordinate, ...]:
|
|
228
|
+
"""Expands coordinates, removing CartesianProducts and Scalars."""
|
|
229
|
+
expanded = []
|
|
230
|
+
for c in coordinates:
|
|
231
|
+
if isinstance(c, CartesianProduct):
|
|
232
|
+
expanded.extend(c.coordinates)
|
|
233
|
+
elif isinstance(c, Scalar):
|
|
234
|
+
pass
|
|
235
|
+
else:
|
|
236
|
+
expanded.append(c)
|
|
237
|
+
return tuple(expanded)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def _consolidate_coordinates(
|
|
241
|
+
*coordinates: Coordinate,
|
|
242
|
+
) -> tuple[Coordinate, ...]:
|
|
243
|
+
"""Consolidates coordinates, removing SelectedAxes when possible."""
|
|
244
|
+
axes = []
|
|
245
|
+
result = []
|
|
246
|
+
|
|
247
|
+
def reset_axes():
|
|
248
|
+
result.extend(axes)
|
|
249
|
+
axes[:] = []
|
|
250
|
+
|
|
251
|
+
def append_axis(c):
|
|
252
|
+
axes.append(c)
|
|
253
|
+
if len(axes) == c.coordinate.ndim:
|
|
254
|
+
# sucessful consolidation
|
|
255
|
+
result.append(c.coordinate)
|
|
256
|
+
axes[:] = []
|
|
257
|
+
|
|
258
|
+
for c in coordinates:
|
|
259
|
+
if isinstance(c, SelectedAxis) and c.axis == 0:
|
|
260
|
+
# new SelectedAxis to consider consolidating
|
|
261
|
+
reset_axes()
|
|
262
|
+
append_axis(c)
|
|
263
|
+
elif (
|
|
264
|
+
isinstance(c, SelectedAxis)
|
|
265
|
+
and axes
|
|
266
|
+
and c.axis == len(axes)
|
|
267
|
+
and c.coordinate == axes[-1].coordinate
|
|
268
|
+
):
|
|
269
|
+
# continued SelectedAxis to consolidate
|
|
270
|
+
append_axis(c)
|
|
271
|
+
else:
|
|
272
|
+
# coordinate cannot be consolidated
|
|
273
|
+
reset_axes()
|
|
274
|
+
result.append(c)
|
|
275
|
+
|
|
276
|
+
reset_axes()
|
|
277
|
+
|
|
278
|
+
return tuple(result)
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def canonicalize(*coordinates: Coordinate) -> tuple[Coordinate, ...]:
|
|
282
|
+
"""Canonicalize coordinates into a minimum equivalent collection."""
|
|
283
|
+
coordinates = _expand_coordinates(*coordinates)
|
|
284
|
+
coordinates = _consolidate_coordinates(*coordinates)
|
|
285
|
+
existing_dims = collections.Counter()
|
|
286
|
+
for c in coordinates:
|
|
287
|
+
existing_dims.update([d for d in c.dims if d is not None])
|
|
288
|
+
repeated_dims = [dim for dim, count in existing_dims.items() if count > 1]
|
|
289
|
+
if repeated_dims:
|
|
290
|
+
raise ValueError(f'coordinates contain {repeated_dims=}')
|
|
291
|
+
return coordinates
|
|
292
|
+
|
|
293
|
+
|
|
294
|
+
T = TypeVar('T')
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
def _concat_tuples(tuples: Iterable[tuple[T, ...]]) -> tuple[T, ...]:
|
|
298
|
+
"""Concatenates tuples."""
|
|
299
|
+
return tuple(itertools.chain(*tuples))
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
K = TypeVar('K')
|
|
303
|
+
V = TypeVar('V')
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
def _merge_dicts(dicts: Iterable[dict[K, V]]) -> dict[K, V]:
|
|
307
|
+
"""Merges dicts."""
|
|
308
|
+
result = {}
|
|
309
|
+
for d in dicts:
|
|
310
|
+
result.update(d)
|
|
311
|
+
return result
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
@utils.export
|
|
315
|
+
@jax.tree_util.register_static
|
|
316
|
+
@dataclasses.dataclass(frozen=True)
|
|
317
|
+
class CartesianProduct(Coordinate):
|
|
318
|
+
"""Coordinate defined as the outer product of independent coordinates."""
|
|
319
|
+
|
|
320
|
+
coordinates: tuple[Coordinate, ...]
|
|
321
|
+
|
|
322
|
+
def __post_init__(self):
|
|
323
|
+
coordinates = canonicalize(*self.coordinates)
|
|
324
|
+
object.__setattr__(self, 'coordinates', coordinates)
|
|
325
|
+
|
|
326
|
+
def __eq__(self, other):
|
|
327
|
+
# TODO(shoyer): require exact equality of coordinate types?
|
|
328
|
+
if not isinstance(other, CartesianProduct):
|
|
329
|
+
return len(self.coordinates) == 1 and self.coordinates[0] == other
|
|
330
|
+
return isinstance(other, CartesianProduct) and self.axes == other.axes
|
|
331
|
+
|
|
332
|
+
@property
|
|
333
|
+
def dims(self) -> tuple[str | None, ...]:
|
|
334
|
+
return _concat_tuples(c.dims for c in self.coordinates)
|
|
335
|
+
|
|
336
|
+
@property
|
|
337
|
+
def shape(self) -> tuple[int, ...]:
|
|
338
|
+
"""Returns the shape of the coordinate axes."""
|
|
339
|
+
return _concat_tuples(c.shape for c in self.coordinates)
|
|
340
|
+
|
|
341
|
+
@property
|
|
342
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
343
|
+
"""Returns a mapping from field names to their values."""
|
|
344
|
+
return _merge_dicts(c.fields for c in self.coordinates)
|
|
345
|
+
|
|
346
|
+
@property
|
|
347
|
+
def axes(self) -> tuple[Coordinate, ...]:
|
|
348
|
+
"""Returns a tuple of Axis objects for each dimension."""
|
|
349
|
+
return _concat_tuples(c.axes for c in self.coordinates)
|
|
350
|
+
|
|
351
|
+
|
|
352
|
+
@utils.export
|
|
353
|
+
@jax.tree_util.register_static
|
|
354
|
+
@dataclasses.dataclass(frozen=True)
|
|
355
|
+
class SizedAxis(Coordinate):
|
|
356
|
+
"""One dimensional coordinate with fixed size but no associated fields."""
|
|
357
|
+
|
|
358
|
+
name: str
|
|
359
|
+
size: int
|
|
360
|
+
|
|
361
|
+
@property
|
|
362
|
+
def dims(self) -> tuple[str, ...]:
|
|
363
|
+
return (self.name,)
|
|
364
|
+
|
|
365
|
+
@property
|
|
366
|
+
def shape(self) -> tuple[int, ...]:
|
|
367
|
+
return (self.size,)
|
|
368
|
+
|
|
369
|
+
@property
|
|
370
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
371
|
+
return {}
|
|
372
|
+
|
|
373
|
+
def __repr__(self):
|
|
374
|
+
return f'coordax.SizedAxis({self.name!r}, size={self.size})'
|
|
375
|
+
|
|
376
|
+
@classmethod
|
|
377
|
+
def from_xarray(
|
|
378
|
+
cls, dims: tuple[str, ...], coords: xarray.Coordinates
|
|
379
|
+
) -> Self | NoCoordinateMatch:
|
|
380
|
+
dim = dims[0]
|
|
381
|
+
if dim in coords:
|
|
382
|
+
return NoCoordinateMatch(
|
|
383
|
+
'can only reconstruct SizedAxis objects from xarray dimensions'
|
|
384
|
+
' without associated coordinate variables, but found a coordinate'
|
|
385
|
+
f' variable for dimension {dim!r}'
|
|
386
|
+
)
|
|
387
|
+
for name, coord in coords.variables.items():
|
|
388
|
+
if dim in coord.dims:
|
|
389
|
+
return NoCoordinateMatch(
|
|
390
|
+
'can only reconstruct SizedAxis objects from xarray dimensions'
|
|
391
|
+
' if the dimensions is not found on any coordinate variables, but '
|
|
392
|
+
f' found a coordinate variable for dimension {dim!r} on {name!r}'
|
|
393
|
+
)
|
|
394
|
+
return cls(dim, size=coords.sizes[dim])
|
|
395
|
+
|
|
396
|
+
|
|
397
|
+
@utils.export
|
|
398
|
+
@dataclasses.dataclass(frozen=True)
|
|
399
|
+
class DummyAxis(Coordinate):
|
|
400
|
+
"""Dummy coordinate for dimensions without associated coordinate values.
|
|
401
|
+
|
|
402
|
+
DummyAxis are placeholders for dimensions that do not have associated
|
|
403
|
+
coordinate values. They are automatically dropped from the Field constructor,
|
|
404
|
+
but are useful for specifying how to construct fields with missing dimension
|
|
405
|
+
names and/or coordinates.
|
|
406
|
+
"""
|
|
407
|
+
|
|
408
|
+
name: str | None
|
|
409
|
+
size: int
|
|
410
|
+
|
|
411
|
+
@property
|
|
412
|
+
def dims(self) -> tuple[str | None, ...]:
|
|
413
|
+
return (self.name,)
|
|
414
|
+
|
|
415
|
+
@property
|
|
416
|
+
def shape(self) -> tuple[int, ...]:
|
|
417
|
+
return (self.size,)
|
|
418
|
+
|
|
419
|
+
@property
|
|
420
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
421
|
+
return {}
|
|
422
|
+
|
|
423
|
+
def __repr__(self):
|
|
424
|
+
return f'coordax.DummyAxis({self.name!r}, size={self.size})'
|
|
425
|
+
|
|
426
|
+
@classmethod
|
|
427
|
+
def from_xarray(
|
|
428
|
+
cls, dims: tuple[str, ...], coords: xarray.Coordinates
|
|
429
|
+
) -> Self | NoCoordinateMatch:
|
|
430
|
+
dim = dims[0]
|
|
431
|
+
for name, coord in coords.variables.items():
|
|
432
|
+
if dim in coord.dims:
|
|
433
|
+
return NoCoordinateMatch(
|
|
434
|
+
f'cannot omit a Coordinate object for dimension {dim!r}'
|
|
435
|
+
f' because it is used by at least one coordinate variable: {name!r}'
|
|
436
|
+
)
|
|
437
|
+
return cls(name=dim, size=coords.sizes[dim])
|
|
438
|
+
|
|
439
|
+
|
|
440
|
+
# TODO(dkochkov): consider storing tuple values instead of np.ndarray (which
|
|
441
|
+
# could be exposed as a property).
|
|
442
|
+
@utils.export
|
|
443
|
+
@jax.tree_util.register_static
|
|
444
|
+
@dataclasses.dataclass(frozen=True)
|
|
445
|
+
class LabeledAxis(Coordinate):
|
|
446
|
+
"""One dimensional coordinate with custom coordinate values."""
|
|
447
|
+
|
|
448
|
+
name: str
|
|
449
|
+
ticks: np.ndarray
|
|
450
|
+
|
|
451
|
+
def __post_init__(self):
|
|
452
|
+
object.__setattr__(self, 'ticks', np.asarray(self.ticks))
|
|
453
|
+
if self.ticks.ndim != 1:
|
|
454
|
+
raise ValueError(f'ticks must be a 1D array, got {self.ticks.shape=}')
|
|
455
|
+
|
|
456
|
+
@property
|
|
457
|
+
def dims(self) -> tuple[str, ...]:
|
|
458
|
+
return (self.name,)
|
|
459
|
+
|
|
460
|
+
@property
|
|
461
|
+
def shape(self) -> tuple[int, ...]:
|
|
462
|
+
return self.ticks.shape
|
|
463
|
+
|
|
464
|
+
@property
|
|
465
|
+
def fields(self) -> dict[str, fields.Field]:
|
|
466
|
+
# needs local import to avoid circular dependency
|
|
467
|
+
from coordax import fields # pylint: disable=g-import-not-at-top
|
|
468
|
+
|
|
469
|
+
return {self.name: fields.wrap(self.ticks, self)}
|
|
470
|
+
|
|
471
|
+
def _components(self):
|
|
472
|
+
return (self.name, ArrayKey(self.ticks))
|
|
473
|
+
|
|
474
|
+
def __eq__(self, other):
|
|
475
|
+
return (
|
|
476
|
+
isinstance(other, LabeledAxis)
|
|
477
|
+
and self._components() == other._components()
|
|
478
|
+
)
|
|
479
|
+
|
|
480
|
+
def __hash__(self) -> int:
|
|
481
|
+
return hash(self._components())
|
|
482
|
+
|
|
483
|
+
def __repr__(self):
|
|
484
|
+
return f'coordax.LabeledAxis({self.name!r}, ticks={self.ticks!r})'
|
|
485
|
+
|
|
486
|
+
@classmethod
|
|
487
|
+
def from_xarray(
|
|
488
|
+
cls, dims: tuple[str, ...], coords: xarray.Coordinates
|
|
489
|
+
) -> Self | NoCoordinateMatch:
|
|
490
|
+
dim = dims[0]
|
|
491
|
+
if dim not in coords:
|
|
492
|
+
return NoCoordinateMatch(
|
|
493
|
+
f'no associated coordinate for dimension {dim!r}'
|
|
494
|
+
)
|
|
495
|
+
if coords[dim].ndim != 1:
|
|
496
|
+
return NoCoordinateMatch(f'coordinate for dimension {dim!r} is not 1D')
|
|
497
|
+
return cls(dim, coords[dim].data)
|
|
498
|
+
|
|
499
|
+
|
|
500
|
+
@utils.export
|
|
501
|
+
def compose(*coordinates: Coordinate) -> Coordinate:
|
|
502
|
+
"""Compose coordinates into a unified coordinate system."""
|
|
503
|
+
product = CartesianProduct(coordinates)
|
|
504
|
+
match len(product.coordinates):
|
|
505
|
+
case 0:
|
|
506
|
+
return Scalar()
|
|
507
|
+
case 1:
|
|
508
|
+
return product.coordinates[0]
|
|
509
|
+
case _:
|
|
510
|
+
return product
|
|
511
|
+
|
|
512
|
+
|
|
513
|
+
def from_xarray(
|
|
514
|
+
data_array: xarray.DataArray,
|
|
515
|
+
coord_types: Sequence[type[Coordinate]] = (LabeledAxis, DummyAxis),
|
|
516
|
+
) -> Coordinate:
|
|
517
|
+
"""Convert the coordinates of an xarray.DataArray into a coordax.Coordinate.
|
|
518
|
+
|
|
519
|
+
Args:
|
|
520
|
+
data_array: xarray.DataArray whose coordinates should be converted.
|
|
521
|
+
coord_types: sequence of coordax.Coordinate subclasses with `from_xarray`
|
|
522
|
+
methods defined. The first coordinate class that returns a coordinate
|
|
523
|
+
object (indicating a match) will be used. By default, coordinates will use
|
|
524
|
+
only generic coordax.LabeledAxis objects. CardesianProduct type is omitted
|
|
525
|
+
from this sequence since it is introduced by the compose() method.
|
|
526
|
+
|
|
527
|
+
Returns:
|
|
528
|
+
A coordax.Coordinate object representing the coordinates of the input
|
|
529
|
+
DataArray.
|
|
530
|
+
"""
|
|
531
|
+
dims = data_array.dims
|
|
532
|
+
coords = []
|
|
533
|
+
|
|
534
|
+
if not all(isinstance(dim, str) for dim in dims):
|
|
535
|
+
raise TypeError(
|
|
536
|
+
'can only convert DataArray objects with string dimensions to Field'
|
|
537
|
+
)
|
|
538
|
+
dims = typing.cast(tuple[str, ...], dims)
|
|
539
|
+
|
|
540
|
+
if not coord_types:
|
|
541
|
+
raise ValueError('coord_types must be non-empty')
|
|
542
|
+
|
|
543
|
+
def get_next_match():
|
|
544
|
+
reasons = []
|
|
545
|
+
for coord_type in coord_types:
|
|
546
|
+
if coord_type == CartesianProduct or coord_type == Scalar:
|
|
547
|
+
continue
|
|
548
|
+
result = coord_type.from_xarray(dims, data_array.coords)
|
|
549
|
+
if isinstance(result, Coordinate):
|
|
550
|
+
return result
|
|
551
|
+
assert isinstance(result, NoCoordinateMatch)
|
|
552
|
+
coord_name = coord_type.__module__ + '.' + coord_type.__name__
|
|
553
|
+
reasons.append(f'{coord_name}: {result.reason}')
|
|
554
|
+
|
|
555
|
+
reasons_str = '\n'.join(reasons)
|
|
556
|
+
raise ValueError(
|
|
557
|
+
'failed to convert xarray.DataArray to coordax.Field, because no '
|
|
558
|
+
f'coordinate type matched the dimensions starting with {dims}:\n'
|
|
559
|
+
f'{data_array}\n\n'
|
|
560
|
+
f'Reasons why coordinate matching failed:\n{reasons_str}'
|
|
561
|
+
)
|
|
562
|
+
|
|
563
|
+
while dims:
|
|
564
|
+
coord = get_next_match()
|
|
565
|
+
coords.append(coord)
|
|
566
|
+
assert coord.ndim > 0 # dimensions will shrink by at least one
|
|
567
|
+
dims = dims[coord.ndim :]
|
|
568
|
+
|
|
569
|
+
return compose(*coords)
|