coordax 0.2.2__tar.gz → 0.2.4__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.2 → coordax-0.2.4}/PKG-INFO +1 -1
  2. {coordax-0.2.2 → coordax-0.2.4}/coordax/__init__.py +1 -1
  3. {coordax-0.2.2 → coordax-0.2.4}/coordax/coordinate_systems.py +48 -7
  4. {coordax-0.2.2 → coordax-0.2.4}/coordax/coordinate_systems_test.py +21 -0
  5. {coordax-0.2.2 → coordax-0.2.4}/coordax/coords.py +1 -0
  6. {coordax-0.2.2 → coordax-0.2.4}/coordax/fields.py +1 -1
  7. {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/PKG-INFO +1 -1
  8. {coordax-0.2.2 → coordax-0.2.4}/pyproject.toml +1 -1
  9. {coordax-0.2.2 → coordax-0.2.4}/LICENSE +0 -0
  10. {coordax-0.2.2 → coordax-0.2.4}/README.md +0 -0
  11. {coordax-0.2.2 → coordax-0.2.4}/coordax/doctest_test.py +0 -0
  12. {coordax-0.2.2 → coordax-0.2.4}/coordax/experimental.py +0 -0
  13. {coordax-0.2.2 → coordax-0.2.4}/coordax/fields_test.py +0 -0
  14. {coordax-0.2.2 → coordax-0.2.4}/coordax/jax_datetime_integration_test.py +0 -0
  15. {coordax-0.2.2 → coordax-0.2.4}/coordax/named_axes.py +0 -0
  16. {coordax-0.2.2 → coordax-0.2.4}/coordax/named_axes_test.py +0 -0
  17. {coordax-0.2.2 → coordax-0.2.4}/coordax/ndarrays.py +0 -0
  18. {coordax-0.2.2 → coordax-0.2.4}/coordax/testing.py +0 -0
  19. {coordax-0.2.2 → coordax-0.2.4}/coordax/utils.py +0 -0
  20. {coordax-0.2.2 → coordax-0.2.4}/coordax/xarray_test.py +0 -0
  21. {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/SOURCES.txt +0 -0
  22. {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/dependency_links.txt +0 -0
  23. {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/requires.txt +0 -0
  24. {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/top_level.txt +0 -0
  25. {coordax-0.2.2 → coordax-0.2.4}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: coordax
3
- Version: 0.2.2
3
+ Version: 0.2.4
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
@@ -53,4 +53,4 @@ from coordax.ndarrays import (
53
53
  )
54
54
  import coordax.testing # pylint: disable=unused-import
55
55
 
56
- __version__ = '0.2.2' # keep sync with pyproject.toml
56
+ __version__ = '0.2.4' # keep sync with pyproject.toml
@@ -25,7 +25,7 @@ import dataclasses
25
25
  import functools
26
26
  import itertools
27
27
  import typing
28
- from typing import Any, Self, TYPE_CHECKING, TypeAlias, TypeVar
28
+ from typing import Any, Self, TYPE_CHECKING, Type, TypeAlias, TypeGuard, TypeVar
29
29
  import warnings
30
30
 
31
31
  from coordax import utils
@@ -35,7 +35,7 @@ import numpy as np
35
35
  if TYPE_CHECKING:
36
36
  # import only under TYPE_CHECKING to avoid circular dependency
37
37
  # pylint: disable=g-bad-import-order
38
- from coordax import fields
38
+ from coordax import fields # pylint: disable=unused-import
39
39
  import xarray
40
40
 
41
41
 
@@ -84,7 +84,7 @@ class Coordinate(abc.ABC):
84
84
  raise NotImplementedError()
85
85
 
86
86
  @property
87
- def fields(self) -> dict[str, fields.Field]:
87
+ def fields(self) -> dict[str, 'fields.Field']:
88
88
  """Optional dict that maps from field names to their values."""
89
89
  return {}
90
90
 
@@ -110,7 +110,7 @@ class Coordinate(abc.ABC):
110
110
 
111
111
  def to_xarray(self) -> dict[str, xarray.Variable]:
112
112
  """Convert this coordinate into xarray variables."""
113
- import xarray
113
+ import xarray # pylint: disable=g-import-not-at-top
114
114
 
115
115
  variables = {}
116
116
  dims_set = {dim for dim in self.dims if dim is not None}
@@ -164,6 +164,12 @@ class ArrayKey:
164
164
  return hash((self.value.shape, self.value.tobytes()))
165
165
 
166
166
 
167
+ @utils.export
168
+ def is_coord(obj: Any) -> TypeGuard[Coordinate]:
169
+ """Returns True if obj is a Coordinate."""
170
+ return isinstance(obj, Coordinate)
171
+
172
+
167
173
  @utils.export
168
174
  @jax.tree_util.register_static
169
175
  @dataclasses.dataclass(frozen=True)
@@ -211,7 +217,7 @@ class SelectedAxis(Coordinate):
211
217
  return (self.coordinate.shape[self.axis],)
212
218
 
213
219
  @property
214
- def fields(self) -> dict[str, fields.Field]:
220
+ def fields(self) -> dict[str, 'fields.Field']:
215
221
  """A maps from field names to their values."""
216
222
  return self.coordinate.fields
217
223
 
@@ -364,7 +370,7 @@ class CartesianProduct(Coordinate):
364
370
  return _concat_tuples(c.shape for c in self.coordinates)
365
371
 
366
372
  @property
367
- def fields(self) -> dict[str, fields.Field]:
373
+ def fields(self) -> dict[str, 'fields.Field']:
368
374
  """Returns a mapping from field names to their values."""
369
375
  return _merge_dicts(c.fields for c in self.coordinates)
370
376
 
@@ -479,7 +485,7 @@ class LabeledAxis(Coordinate):
479
485
  return self.ticks.shape
480
486
 
481
487
  @property
482
- def fields(self) -> dict[str, fields.Field]:
488
+ def fields(self) -> dict[str, 'fields.Field']:
483
489
  # needs local import to avoid circular dependency
484
490
  from coordax import fields # pylint: disable=g-import-not-at-top
485
491
 
@@ -636,6 +642,41 @@ def replace_axes(
636
642
  return compose(*axes)
637
643
 
638
644
 
645
+ @functools.partial(utils.export, module='coordax.coords')
646
+ def extract(
647
+ coord: Coordinate,
648
+ component_type: Type[Coordinate] | tuple[Type[Coordinate], ...],
649
+ ) -> Coordinate:
650
+ """Extracts component of type `component_type` from the `coord`.
651
+
652
+ Args:
653
+ coord: The coordinate system to search.
654
+ component_type: The type(s) of coordinate to extract.
655
+
656
+ Returns:
657
+ The single coordinate component matching the given type(s).
658
+
659
+ Raises:
660
+ ValueError: If there is not exactly one component of the given type(s).
661
+
662
+ Examples:
663
+ >>> import coordax as cx
664
+ >>> import numpy as np
665
+ >>> x = cx.LabeledAxis('x', np.linspace(0, np.pi, 4))
666
+ >>> y = cx.SizedAxis('y', 3)
667
+ >>> cx.coords.extract(cx.coords.compose(x, y), cx.SizedAxis)
668
+ coordax.SizedAxis('y', size=3)
669
+ """
670
+ components = canonicalize(coord)
671
+ of_type = [c for c in components if isinstance(c, component_type)]
672
+ if len(of_type) != 1:
673
+ raise ValueError(
674
+ f'Expected exactly one instance of {component_type}, found {of_type}'
675
+ )
676
+ [result] = of_type
677
+ return result
678
+
679
+
639
680
  @functools.partial(utils.export, module='coordax.coords')
640
681
  def from_xarray(
641
682
  data_array: xarray.DataArray,
@@ -459,6 +459,27 @@ class CoordinateSystemsTest(parameterized.TestCase):
459
459
  ):
460
460
  cx.CartesianProduct((x, y))
461
461
 
462
+ def test_extract(self):
463
+ x = cx.SizedAxis('x', 2)
464
+ y = cx.SizedAxis('y', 3)
465
+ z = cx.LabeledAxis('z', np.arange(4))
466
+ xy = cx.coords.compose(x, y)
467
+ xz = cx.coords.compose(x, z)
468
+
469
+ self.assertEqual(cx.coords.extract(x, cx.SizedAxis), x)
470
+ self.assertEqual(cx.coords.extract(z, cx.LabeledAxis), z)
471
+ self.assertEqual(cx.coords.extract(xz, cx.LabeledAxis), z)
472
+ self.assertEqual(cx.coords.extract(z, (cx.SizedAxis, cx.LabeledAxis)), z)
473
+
474
+ with self.assertRaisesRegex(ValueError, 'Expected exactly one instance'):
475
+ cx.coords.extract(xy, cx.SizedAxis)
476
+
477
+ with self.assertRaisesRegex(ValueError, 'Expected exactly one instance'):
478
+ cx.coords.extract(xz, (cx.SizedAxis, cx.LabeledAxis))
479
+
480
+ with self.assertRaisesRegex(ValueError, 'Expected exactly one instance'):
481
+ cx.coords.extract(x, cx.LabeledAxis)
482
+
462
483
  def test_deprecated_aliases(self):
463
484
  with self.assertWarnsRegex(
464
485
  DeprecationWarning,
@@ -22,6 +22,7 @@ from coordax.coordinate_systems import (
22
22
  insert_axes as insert_axes,
23
23
  replace_axes as replace_axes,
24
24
  from_xarray as from_xarray,
25
+ extract as extract,
25
26
  ArrayKey as ArrayKey,
26
27
  NoCoordinateMatch as NoCoordinateMatch,
27
28
  )
@@ -538,7 +538,7 @@ class Field:
538
538
  [0., 0., 0.]], dtype=float32)
539
539
  Dimensions without coordinates: x, y
540
540
  """
541
- import xarray
541
+ import xarray # pylint: disable=g-import-not-at-top
542
542
 
543
543
  if not all(isinstance(dim, str) for dim in self.dims):
544
544
  raise ValueError(
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: coordax
3
- Version: 0.2.2
3
+ Version: 0.2.4
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
@@ -7,7 +7,7 @@ packages = ["coordax"]
7
7
 
8
8
  [project]
9
9
  name = "coordax"
10
- version = "0.2.2" # keep sync with __init__.py
10
+ version = "0.2.4" # keep sync with __init__.py
11
11
  description = "Coordinate axes for scientific computing in JAX"
12
12
  authors = [
13
13
  {name = "Google LLC", email = "noreply@google.com"},
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes