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 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)