coordax 0.2.0__tar.gz → 0.2.2__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {coordax-0.2.0 → coordax-0.2.2}/PKG-INFO +1 -1
- {coordax-0.2.0 → coordax-0.2.2}/coordax/__init__.py +2 -1
- {coordax-0.2.0 → coordax-0.2.2}/coordax/coordinate_systems.py +120 -28
- coordax-0.2.2/coordax/doctest_test.py +31 -0
- coordax-0.2.2/coordax/experimental.py +22 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax/fields.py +460 -69
- {coordax-0.2.0 → coordax-0.2.2}/coordax/jax_datetime_integration_test.py +1 -1
- {coordax-0.2.0 → coordax-0.2.2}/coordax/named_axes.py +32 -32
- {coordax-0.2.0 → coordax-0.2.2}/coordax/named_axes_test.py +6 -6
- {coordax-0.2.0 → coordax-0.2.2}/coordax/ndarrays.py +21 -7
- {coordax-0.2.0 → coordax-0.2.2}/coordax/testing.py +3 -3
- {coordax-0.2.0 → coordax-0.2.2}/coordax/xarray_test.py +20 -7
- {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/PKG-INFO +1 -1
- {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/SOURCES.txt +2 -0
- {coordax-0.2.0 → coordax-0.2.2}/pyproject.toml +1 -1
- {coordax-0.2.0 → coordax-0.2.2}/LICENSE +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/README.md +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax/coordinate_systems_test.py +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax/coords.py +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax/fields_test.py +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax/utils.py +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/dependency_links.txt +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/requires.txt +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/top_level.txt +0 -0
- {coordax-0.2.0 → coordax-0.2.2}/setup.cfg +0 -0
|
@@ -41,6 +41,7 @@ from coordax.fields import (
|
|
|
41
41
|
cmap as cmap,
|
|
42
42
|
cpmap as cpmap,
|
|
43
43
|
get_coordinate as get_coordinate,
|
|
44
|
+
from_xarray as from_xarray,
|
|
44
45
|
wrap_like as wrap_like, # deprecated
|
|
45
46
|
wrap as wrap, # deprecated
|
|
46
47
|
tag as tag,
|
|
@@ -52,4 +53,4 @@ from coordax.ndarrays import (
|
|
|
52
53
|
)
|
|
53
54
|
import coordax.testing # pylint: disable=unused-import
|
|
54
55
|
|
|
55
|
-
__version__ = '0.2.
|
|
56
|
+
__version__ = '0.2.2' # keep sync with pyproject.toml
|
|
@@ -46,7 +46,7 @@ Sequence = collections.abc.Sequence
|
|
|
46
46
|
@functools.partial(utils.export, module='coordax.coords')
|
|
47
47
|
@dataclasses.dataclass(frozen=True)
|
|
48
48
|
class NoCoordinateMatch:
|
|
49
|
-
"""For use when
|
|
49
|
+
"""For use when no Coordax coordinate matches xarray coordinate."""
|
|
50
50
|
|
|
51
51
|
reason: str
|
|
52
52
|
|
|
@@ -56,18 +56,15 @@ class Coordinate(abc.ABC):
|
|
|
56
56
|
"""Abstract class for coordinate objects.
|
|
57
57
|
|
|
58
58
|
Coordinate subclasses are expected to obey several invariants:
|
|
59
|
-
1. Dimension names may not be repeated:
|
|
60
|
-
2. All dimensions must be named:
|
|
59
|
+
1. Dimension names may not be repeated: ``len(set(dims)) == len(dims)``
|
|
60
|
+
2. All dimensions must be named: ``len(shape) == len(dims)``
|
|
61
61
|
|
|
62
62
|
Every non-abstract Coordinate subclass must be registered as a "static"
|
|
63
63
|
pytree node, e.g., by decorating the class with
|
|
64
|
-
|
|
65
|
-
|
|
64
|
+
``@jax.tree_util.register_static``. Static pytrees nodes must implement
|
|
65
|
+
``__hash__`` and ``__eq__`` according to the requirements of keys in Python
|
|
66
66
|
dictionaries. This is easiest to acheive with frozen dataclasses, but care
|
|
67
67
|
must be taken when working with np.ndarray attributes.
|
|
68
|
-
|
|
69
|
-
TODO(shoyer): add documentation examples, including a version using ArrayKey
|
|
70
|
-
to wrap np.ndarray attributions.
|
|
71
68
|
"""
|
|
72
69
|
|
|
73
70
|
@property
|
|
@@ -76,7 +73,7 @@ class Coordinate(abc.ABC):
|
|
|
76
73
|
"""Dimension names of the coordinate.
|
|
77
74
|
|
|
78
75
|
All subclasses must return a tuple of dimension names as strings, with the
|
|
79
|
-
exception of
|
|
76
|
+
exception of ``DummyAxis``.
|
|
80
77
|
"""
|
|
81
78
|
raise NotImplementedError()
|
|
82
79
|
|
|
@@ -138,8 +135,12 @@ class Coordinate(abc.ABC):
|
|
|
138
135
|
values.
|
|
139
136
|
|
|
140
137
|
Returns:
|
|
141
|
-
A matching instance of this coordinate or
|
|
142
|
-
coordinate does not match the xarray dimensions and coordinates.
|
|
138
|
+
A matching instance of this coordinate or ``NoCoordinateMatch`` if this
|
|
139
|
+
coordinate type does not match the xarray dimensions and coordinates.
|
|
140
|
+
|
|
141
|
+
See also:
|
|
142
|
+
:func:`coordax.from_xarray`
|
|
143
|
+
:func:`coordax.coords.from_xarray`
|
|
143
144
|
"""
|
|
144
145
|
raise NotImplementedError('from_xarray not implemented')
|
|
145
146
|
|
|
@@ -167,7 +168,7 @@ class ArrayKey:
|
|
|
167
168
|
@jax.tree_util.register_static
|
|
168
169
|
@dataclasses.dataclass(frozen=True)
|
|
169
170
|
class Scalar(Coordinate):
|
|
170
|
-
"""Zero dimensional sentinel coordinate used to label
|
|
171
|
+
"""Zero dimensional sentinel coordinate used to label scalar fields."""
|
|
171
172
|
|
|
172
173
|
@property
|
|
173
174
|
def dims(self) -> tuple[str, ...]:
|
|
@@ -280,7 +281,27 @@ def _consolidate_coordinates(
|
|
|
280
281
|
|
|
281
282
|
@functools.partial(utils.export, module='coordax.coords')
|
|
282
283
|
def canonicalize(*coordinates: Coordinate) -> tuple[Coordinate, ...]:
|
|
283
|
-
"""Canonicalize coordinates into a minimum equivalent collection.
|
|
284
|
+
"""Canonicalize coordinates into a minimum equivalent collection.
|
|
285
|
+
|
|
286
|
+
Args:
|
|
287
|
+
*coordinates: The coordinates to canonicalize.
|
|
288
|
+
|
|
289
|
+
Returns:
|
|
290
|
+
A tuple of canonicalized coordinates, where ``CartesianProduct`` objects are
|
|
291
|
+
flattened, ``Scalar`` objects are removed and ``SelectedAxis`` objects are
|
|
292
|
+
merged back into the original coordinate if all axes are selected in order.
|
|
293
|
+
|
|
294
|
+
Examples:
|
|
295
|
+
>>> import coordax as cx
|
|
296
|
+
>>> x = cx.SizedAxis('x', 2)
|
|
297
|
+
>>> y = cx.SizedAxis('y', 3)
|
|
298
|
+
>>> cx.coords.canonicalize(x, y)
|
|
299
|
+
(coordax.SizedAxis('x', size=2), coordax.SizedAxis('y', size=3))
|
|
300
|
+
|
|
301
|
+
>>> xy = cx.coords.compose(x, y)
|
|
302
|
+
>>> cx.coords.canonicalize(xy)
|
|
303
|
+
(coordax.SizedAxis('x', size=2), coordax.SizedAxis('y', size=3))
|
|
304
|
+
"""
|
|
284
305
|
coordinates = _expand_coordinates(*coordinates)
|
|
285
306
|
coordinates = _consolidate_coordinates(*coordinates)
|
|
286
307
|
existing_dims = collections.Counter()
|
|
@@ -316,7 +337,10 @@ def _merge_dicts(dicts: Iterable[dict[K, V]]) -> dict[K, V]:
|
|
|
316
337
|
@jax.tree_util.register_static
|
|
317
338
|
@dataclasses.dataclass(frozen=True)
|
|
318
339
|
class CartesianProduct(Coordinate):
|
|
319
|
-
"""Coordinate defined as the outer product of independent coordinates.
|
|
340
|
+
"""Coordinate defined as the outer product of independent coordinates.
|
|
341
|
+
|
|
342
|
+
To construct a ``CartesianProduct``, use :func:`coordax.coords.compose`.
|
|
343
|
+
"""
|
|
320
344
|
|
|
321
345
|
coordinates: tuple[Coordinate, ...]
|
|
322
346
|
|
|
@@ -492,7 +516,24 @@ class LabeledAxis(Coordinate):
|
|
|
492
516
|
|
|
493
517
|
@functools.partial(utils.export, module='coordax.coords')
|
|
494
518
|
def compose(*coordinates: Coordinate) -> Coordinate:
|
|
495
|
-
|
|
519
|
+
# pylint: disable=line-too-long
|
|
520
|
+
# fmt: off
|
|
521
|
+
"""Compose coordinates into a unified coordinate system.
|
|
522
|
+
|
|
523
|
+
Args:
|
|
524
|
+
*coordinates: The coordinates to compose.
|
|
525
|
+
|
|
526
|
+
Returns:
|
|
527
|
+
A single coordinate object representing the Cartesian product of the inputs.
|
|
528
|
+
|
|
529
|
+
Examples:
|
|
530
|
+
>>> import coordax as cx
|
|
531
|
+
>>> x = cx.SizedAxis('x', 2)
|
|
532
|
+
>>> y = cx.SizedAxis('y', 3)
|
|
533
|
+
>>> cx.coords.compose(x, y)
|
|
534
|
+
CartesianProduct(coordinates=(coordax.SizedAxis('x', size=2), coordax.SizedAxis('y', size=3)))
|
|
535
|
+
"""
|
|
536
|
+
# fmt: on
|
|
496
537
|
product = CartesianProduct(coordinates)
|
|
497
538
|
match len(product.coordinates):
|
|
498
539
|
case 0:
|
|
@@ -508,7 +549,26 @@ def insert_axes(
|
|
|
508
549
|
coordinate: Coordinate,
|
|
509
550
|
indices_to_axes: dict[int, Coordinate],
|
|
510
551
|
) -> Coordinate:
|
|
511
|
-
|
|
552
|
+
# pylint: disable=line-too-long
|
|
553
|
+
# fmt: off
|
|
554
|
+
"""Returns ``coordinate`` with extra axes inserted at specified positions.
|
|
555
|
+
|
|
556
|
+
Args:
|
|
557
|
+
coordinate: The coordinate system to modify.
|
|
558
|
+
indices_to_axes: A mapping from insertion index to the new coordinate to
|
|
559
|
+
insert. Indices are relative to the *output* coordinate system.
|
|
560
|
+
|
|
561
|
+
Returns:
|
|
562
|
+
A new coordinate object with the axes inserted.
|
|
563
|
+
|
|
564
|
+
Examples:
|
|
565
|
+
>>> import coordax as cx
|
|
566
|
+
>>> x = cx.SizedAxis('x', 2)
|
|
567
|
+
>>> z = cx.SizedAxis('z', 4)
|
|
568
|
+
>>> cx.coords.insert_axes(x, {1: z})
|
|
569
|
+
CartesianProduct(coordinates=(coordax.SizedAxis('x', size=2), coordax.SizedAxis('z', size=4)))
|
|
570
|
+
"""
|
|
571
|
+
# fmt: on
|
|
512
572
|
indices_to_axes = indices_to_axes.copy()
|
|
513
573
|
ndim = coordinate.ndim + len(indices_to_axes)
|
|
514
574
|
normalize_idx = lambda i: i + ndim if i < 0 else i
|
|
@@ -528,7 +588,9 @@ def replace_axes(
|
|
|
528
588
|
to_replace: Coordinate,
|
|
529
589
|
replace_with: Coordinate,
|
|
530
590
|
) -> Coordinate:
|
|
531
|
-
|
|
591
|
+
# pylint: disable=line-too-long
|
|
592
|
+
# fmt: off
|
|
593
|
+
"""Returns ``coordinate`` with ``to_replace`` replaced by ``replace_with``.
|
|
532
594
|
|
|
533
595
|
Args:
|
|
534
596
|
coordinate: The coordinate system to modify.
|
|
@@ -536,11 +598,23 @@ def replace_axes(
|
|
|
536
598
|
replace_with: The new coordinate to insert.
|
|
537
599
|
|
|
538
600
|
Returns:
|
|
539
|
-
A new coordinate object with
|
|
601
|
+
A new coordinate object with ``to_replace`` substituted by ``replace_with``.
|
|
540
602
|
|
|
541
603
|
Raises:
|
|
542
|
-
ValueError: If
|
|
604
|
+
ValueError: If ``to_replace`` is not a contiguous part of ``coordinate``.
|
|
605
|
+
|
|
606
|
+
Examples:
|
|
607
|
+
>>> import coordax as cx
|
|
608
|
+
>>> x = cx.SizedAxis('x', 2)
|
|
609
|
+
>>> y = cx.SizedAxis('y', 3)
|
|
610
|
+
>>> xy = cx.coords.compose(x, y)
|
|
611
|
+
>>> xy
|
|
612
|
+
CartesianProduct(coordinates=(coordax.SizedAxis('x', size=2), coordax.SizedAxis('y', size=3)))
|
|
613
|
+
>>> z = cx.SizedAxis('z', 4)
|
|
614
|
+
>>> cx.coords.replace_axes(xy, x, z)
|
|
615
|
+
CartesianProduct(coordinates=(coordax.SizedAxis('z', size=4), coordax.SizedAxis('y', size=3)))
|
|
543
616
|
"""
|
|
617
|
+
# fmt: on
|
|
544
618
|
if not to_replace.dims:
|
|
545
619
|
raise ValueError(f'`to_replace` must have dimensions, got {to_replace}')
|
|
546
620
|
|
|
@@ -567,20 +641,38 @@ def from_xarray(
|
|
|
567
641
|
data_array: xarray.DataArray,
|
|
568
642
|
coord_types: Sequence[type[Coordinate]] = (LabeledAxis, DummyAxis),
|
|
569
643
|
) -> Coordinate:
|
|
570
|
-
|
|
644
|
+
# pylint: disable=line-too-long
|
|
645
|
+
# fmt: off
|
|
646
|
+
"""Convert coordinates of an ``xarray.DataArray`` into ``coordax.Coordinate``.
|
|
571
647
|
|
|
572
648
|
Args:
|
|
573
|
-
data_array: xarray.DataArray whose coordinates should be converted.
|
|
574
|
-
coord_types: sequence of coordax.Coordinate subclasses with
|
|
575
|
-
methods defined. The first coordinate class that returns a
|
|
576
|
-
object (indicating a match) will be used. By default,
|
|
577
|
-
only generic
|
|
578
|
-
|
|
649
|
+
data_array: ``xarray.DataArray`` whose coordinates should be converted.
|
|
650
|
+
coord_types: sequence of ``coordax.Coordinate`` subclasses with
|
|
651
|
+
``from_xarray`` methods defined. The first coordinate class that returns a
|
|
652
|
+
coordinate object (indicating a match) will be used. By default,
|
|
653
|
+
coordinates will use only generic ``LabeledAxis`` and ``DummyAxis``
|
|
654
|
+
objects.
|
|
579
655
|
|
|
580
656
|
Returns:
|
|
581
|
-
A coordax.Coordinate object
|
|
582
|
-
|
|
657
|
+
A coordax.Coordinate object, constructed from any of the indicated types
|
|
658
|
+
(plus ``CartesianProduct`` and ``Scalar``), representing the coordinates of
|
|
659
|
+
the input DataArray.
|
|
660
|
+
|
|
661
|
+
Raises:
|
|
662
|
+
ValueError: if no matching coordinate is found.
|
|
663
|
+
|
|
664
|
+
Examples:
|
|
665
|
+
>>> import coordax as cx
|
|
666
|
+
>>> import xarray as xr
|
|
667
|
+
>>> import numpy as np
|
|
668
|
+
>>> da = xr.DataArray(np.zeros((2, 3)), dims=('x', 'y'), coords={'x': [1, 2]})
|
|
669
|
+
>>> cx.coords.from_xarray(da)
|
|
670
|
+
CartesianProduct(coordinates=(coordax.LabeledAxis('x', ticks=array([1, 2])), coordax.DummyAxis('y', size=3)))
|
|
671
|
+
|
|
672
|
+
See also:
|
|
673
|
+
:func:`coordax.from_xarray`
|
|
583
674
|
"""
|
|
675
|
+
# fmt: on
|
|
584
676
|
dims = data_array.dims
|
|
585
677
|
coords = []
|
|
586
678
|
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
# Copyright 2025 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
|
+
import doctest
|
|
15
|
+
from absl.testing import absltest
|
|
16
|
+
import coordax
|
|
17
|
+
import pytest
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
pytest.importorskip('xarray')
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def load_tests(loader, tests, ignore):
|
|
24
|
+
tests.addTests(doctest.DocTestSuite(coordax))
|
|
25
|
+
tests.addTests(doctest.DocTestSuite(coordax.coords))
|
|
26
|
+
tests.addTests(doctest.DocTestSuite(coordax.testing))
|
|
27
|
+
return tests
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
if __name__ == '__main__':
|
|
31
|
+
absltest.main()
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
# Copyright 2025 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
|
+
"""Experimental parts of Coordax."""
|
|
15
|
+
|
|
16
|
+
# Note: import <name> as <name> is required for names to be exported.
|
|
17
|
+
# See PEP 484 & https://github.com/jax-ml/jax/issues/7570
|
|
18
|
+
# pylint: disable=g-multiple-import,useless-import-alias,g-importing-member,unused-import
|
|
19
|
+
from coordax.ndarrays import (
|
|
20
|
+
NDArray as NDArray,
|
|
21
|
+
register_ndarray as register_ndarray,
|
|
22
|
+
)
|