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.
Files changed (25) hide show
  1. {coordax-0.2.0 → coordax-0.2.2}/PKG-INFO +1 -1
  2. {coordax-0.2.0 → coordax-0.2.2}/coordax/__init__.py +2 -1
  3. {coordax-0.2.0 → coordax-0.2.2}/coordax/coordinate_systems.py +120 -28
  4. coordax-0.2.2/coordax/doctest_test.py +31 -0
  5. coordax-0.2.2/coordax/experimental.py +22 -0
  6. {coordax-0.2.0 → coordax-0.2.2}/coordax/fields.py +460 -69
  7. {coordax-0.2.0 → coordax-0.2.2}/coordax/jax_datetime_integration_test.py +1 -1
  8. {coordax-0.2.0 → coordax-0.2.2}/coordax/named_axes.py +32 -32
  9. {coordax-0.2.0 → coordax-0.2.2}/coordax/named_axes_test.py +6 -6
  10. {coordax-0.2.0 → coordax-0.2.2}/coordax/ndarrays.py +21 -7
  11. {coordax-0.2.0 → coordax-0.2.2}/coordax/testing.py +3 -3
  12. {coordax-0.2.0 → coordax-0.2.2}/coordax/xarray_test.py +20 -7
  13. {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/PKG-INFO +1 -1
  14. {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/SOURCES.txt +2 -0
  15. {coordax-0.2.0 → coordax-0.2.2}/pyproject.toml +1 -1
  16. {coordax-0.2.0 → coordax-0.2.2}/LICENSE +0 -0
  17. {coordax-0.2.0 → coordax-0.2.2}/README.md +0 -0
  18. {coordax-0.2.0 → coordax-0.2.2}/coordax/coordinate_systems_test.py +0 -0
  19. {coordax-0.2.0 → coordax-0.2.2}/coordax/coords.py +0 -0
  20. {coordax-0.2.0 → coordax-0.2.2}/coordax/fields_test.py +0 -0
  21. {coordax-0.2.0 → coordax-0.2.2}/coordax/utils.py +0 -0
  22. {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/dependency_links.txt +0 -0
  23. {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/requires.txt +0 -0
  24. {coordax-0.2.0 → coordax-0.2.2}/coordax.egg-info/top_level.txt +0 -0
  25. {coordax-0.2.0 → coordax-0.2.2}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: coordax
3
- Version: 0.2.0
3
+ Version: 0.2.2
4
4
  Summary: Coordinate axes for scientific computing in JAX
5
5
  Author-email: Google LLC <noreply@google.com>
6
6
  License-Expression: Apache-2.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.0' # keep sync with pyproject.toml
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 a coordinate does not match an xarray.Coordinates object."""
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: `len(set(dims)) == len(dims)`
60
- 2. All dimensions must be named: `len(shape) == len(dims)`
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
- `@jax.tree_util.register_static`. Static pytrees nodes must implement
65
- `__hash__` and `__eq__` according to the requirements of keys in Python
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 `DummyAxis`.
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 `NoCoordinateMatch` if this
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 stand alone scalars."""
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
- """Compose coordinates into a unified coordinate system."""
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
- """Returns `coordinate` with extra axes inserted at specified positions."""
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
- """Returns `coordinate` with `to_replace` replaced by `replace_with`.
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 `to_replace` substituted by `replace_with`.
601
+ A new coordinate object with ``to_replace`` substituted by ``replace_with``.
540
602
 
541
603
  Raises:
542
- ValueError: If `to_replace` is not a contiguous part of `coordinate`.
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
- """Convert the coordinates of an xarray.DataArray into a coordax.Coordinate.
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 `from_xarray`
575
- methods defined. The first coordinate class that returns a coordinate
576
- object (indicating a match) will be used. By default, coordinates will use
577
- only generic coordax.LabeledAxis objects. CardesianProduct type is omitted
578
- from this sequence since it is introduced by the compose() method.
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 representing the coordinates of the input
582
- DataArray.
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
+ )