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.
- {coordax-0.2.2 → coordax-0.2.4}/PKG-INFO +1 -1
- {coordax-0.2.2 → coordax-0.2.4}/coordax/__init__.py +1 -1
- {coordax-0.2.2 → coordax-0.2.4}/coordax/coordinate_systems.py +48 -7
- {coordax-0.2.2 → coordax-0.2.4}/coordax/coordinate_systems_test.py +21 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/coords.py +1 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/fields.py +1 -1
- {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/PKG-INFO +1 -1
- {coordax-0.2.2 → coordax-0.2.4}/pyproject.toml +1 -1
- {coordax-0.2.2 → coordax-0.2.4}/LICENSE +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/README.md +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/doctest_test.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/experimental.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/fields_test.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/jax_datetime_integration_test.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/named_axes.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/named_axes_test.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/ndarrays.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/testing.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/utils.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax/xarray_test.py +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/SOURCES.txt +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/dependency_links.txt +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/requires.txt +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/coordax.egg-info/top_level.txt +0 -0
- {coordax-0.2.2 → coordax-0.2.4}/setup.cfg +0 -0
|
@@ -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,
|
|
@@ -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(
|
|
@@ -7,7 +7,7 @@ packages = ["coordax"]
|
|
|
7
7
|
|
|
8
8
|
[project]
|
|
9
9
|
name = "coordax"
|
|
10
|
-
version = "0.2.
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|