np-struct 0.1.6__tar.gz → 0.2.0__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: np-struct
3
- Version: 0.1.6
3
+ Version: 0.2.0
4
4
  Summary: Interface for NumPy structured arrays
5
5
  Author-email: Rick Lyon <rlyon14@gmail.com>
6
6
  Project-URL: repository, https://github.com/ricklyon/np_struct
@@ -1,24 +1,17 @@
1
1
  import numpy as np
2
2
  import datetime as dt
3
- from scipy import interpolate, ndimage
4
- from scipy.interpolate import interp1d
3
+ from scipy import ndimage
4
+ from scipy import interpolate
5
5
  from collections import OrderedDict
6
6
  from copy import deepcopy as dcopy
7
7
  import datetime
8
- from itertools import chain
9
- from typing import TYPE_CHECKING
8
+ from typing import TYPE_CHECKING, Callable
9
+ from itertools import product
10
10
 
11
- def check_shapes(a: tuple, b: tuple):
12
- """
13
- Check that the shape tuples a and b match
14
- """
15
-
16
- if len(a) != len(b):
17
- return False
18
-
19
- # check that the length of each dimension matches
20
- return all([a[i] == b[i] for i in range(len(a))])
11
+ from np_struct import utils
21
12
 
13
+ if TYPE_CHECKING:
14
+ from matplotlib import axes
22
15
 
23
16
  def datetime_idx_handler(v: datetime.datetime, coords: np.ndarray):
24
17
  """
@@ -350,12 +343,19 @@ class ldarray(np.ndarray):
350
343
 
351
344
  def __new__(cls, data=None, coords=None, attrs= dict(), dtype=None):
352
345
 
346
+ # if coords is not provided, attempt to use coords from data
347
+ if coords is None and isinstance(data, ldarray):
348
+ coords = data.coords
349
+
353
350
  # cast coords as a OrderedDictionary type
354
- if not isinstance(coords, Coords):
351
+ if coords is not None and not isinstance(coords, Coords):
355
352
  coords = Coords(**coords)
356
353
 
357
354
  # create 0 filled array if no data is given in the constructor
358
355
  if data is None:
356
+ if coords is None:
357
+ raise ValueError("Coords must be provided.")
358
+
359
359
  obj = np.zeros(coords.shape, dtype=dtype).view(cls)
360
360
 
361
361
  # cast input data to ldarray type
@@ -368,8 +368,8 @@ class ldarray(np.ndarray):
368
368
  obj = obj.view(cls)
369
369
 
370
370
  # If dim is not compatible with the data shape return a standard numpy array
371
- if (coords is None) or (not check_shapes(obj.shape, coords.shape)):
372
- raise TypeError(
371
+ if (coords is not None) and (not utils.check_shapes(obj.shape, coords.shape)):
372
+ raise ValueError(
373
373
  "Coordinates of shape {} are not compatible with data of shape {}.".format(coords.shape, obj.shape)
374
374
  )
375
375
 
@@ -419,8 +419,8 @@ class ldarray(np.ndarray):
419
419
  obj = super().__array_function__(func, types, args, kwargs)
420
420
 
421
421
  # recast as ldarray if return type is ndarray
422
- if isinstance(obj, np.ndarray) and not isinstance(obj, ldarray) and check_shapes(obj.shape, self.shape):
423
- obj = ldarray(obj, coords=self.coords)
422
+ if isinstance(obj, np.ndarray) and not isinstance(obj, ldarray) and utils.check_shapes(obj.shape, self.shape):
423
+ obj = ldarray(obj, coords=self.coords, attrs=self.attrs)
424
424
 
425
425
  # remove axis from coords if it was indexed out by np.sum, np.average, etc...
426
426
  elif isinstance(obj, np.ndarray) and axis is not None and len(obj.shape) == (len(self.shape) - len(axis)):
@@ -429,7 +429,7 @@ class ldarray(np.ndarray):
429
429
  keys = tuple(self.coords.keys())
430
430
  [coords.pop(keys[i]) for i in axis]
431
431
  # recast as labeled array
432
- obj = ldarray(obj, coords=coords)
432
+ obj = ldarray(obj, coords=coords, attrs=self.attrs)
433
433
 
434
434
  # invalidate coords for functions capable of leaving the shape intact but permuting the axis
435
435
  # order. transpose is subclassed separately and not included here.
@@ -448,8 +448,9 @@ class ldarray(np.ndarray):
448
448
  # array finalize is called when array is cast to a new type, indexed, or whenever a new array with a different
449
449
  # shape is created (i.e. transpose). By default, drop the coordinates which are most likely out of date now.
450
450
  # Coordinates will be added back by lower level functions if the shape stayed the same.
451
- if isinstance(obj, ldarray) and getattr(obj, "coords", None) and check_shapes(self.shape, obj.coords.shape):
451
+ if isinstance(obj, ldarray) and getattr(obj, "coords", None) and utils.check_shapes(self.shape, obj.coords.shape):
452
452
  self.coords = dcopy(obj.coords)
453
+ self.attrs = dcopy(obj.attrs)
453
454
  else:
454
455
  self.coords = None
455
456
 
@@ -533,13 +534,14 @@ class ldarray(np.ndarray):
533
534
 
534
535
  # if the shapes of the inputs were expanded, restore the full expanded coordinates if the shape
535
536
  # is still consistent.
536
- elif len(result_coords) and check_shapes(results.shape, Coords(**result_coords).shape):
537
- results = ldarray(results, coords=result_coords)
537
+ elif len(result_coords) and utils.check_shapes(results.shape, Coords(**result_coords).shape):
538
+ results = ldarray(results, coords=result_coords, attrs=self.attrs)
538
539
 
539
540
  # if the shape is the same after the math operation, restore the coordinates
540
- elif self.coords and check_shapes(results.shape, self.coords.shape):
541
+ elif self.coords and utils.check_shapes(results.shape, self.coords.shape):
541
542
  results = results.view(ldarray)
542
543
  results.coords = dcopy(self.coords)
544
+ results.attrs = dcopy(self.attrs)
543
545
 
544
546
  else:
545
547
  results = results.view(np.ndarray)
@@ -606,7 +608,16 @@ class ldarray(np.ndarray):
606
608
  # index is a standard index of slices or integers so pass key to the numpy indexing routine.
607
609
  # this object will have the coords set to None by __array_finalize___
608
610
  obj = super(ldarray, self).__getitem__(key)
611
+
612
+ # Cast index key as a tuple if it's a single value
613
+ nkey = tuple(key) if isinstance(key, (tuple, list)) else (key,)
609
614
 
615
+ # use coords from advanced index
616
+ labeled_idx = [k for k in nkey if isinstance(k, ldarray) and utils.check_shapes(k.shape, obj.shape)]
617
+ if len(labeled_idx):
618
+ obj.coords = dcopy(labeled_idx[0].coords)
619
+ return obj
620
+
610
621
  # shape length can be greater after indexing if np.newaxis was used. In this case just
611
622
  # return a standard numpy array and make the user responsible for adding dimensional labels.
612
623
  if len(obj.shape) > len(self.shape):
@@ -623,13 +634,11 @@ class ldarray(np.ndarray):
623
634
  # At this point, we need to index the dimension dictionary so it matches the obj data,
624
635
  # and remove axis that were indexed out completely.
625
636
  try:
626
- # Cast index key as a tuple if it's a single value
627
- nkey = tuple(key) if isinstance(key, (tuple, list)) else (key,)
628
637
 
629
638
  # Initialize list of indices for each dimension that will be used to index the label arrays in dim.
630
639
  # Length is the original array shape length so it matches ndim.
631
640
  idx = [slice(None,None) for i in range(len(self.shape))]
632
-
641
+
633
642
  # step through index keys and update idx with the appropriate keys.
634
643
  # Keys are always in order of the array dimensions, but axis can be skipped with the Ellipsis operator.
635
644
  idx_i = 0
@@ -656,14 +665,15 @@ class ldarray(np.ndarray):
656
665
  # idx has a value for every dimension so we can use i to get the correct index key
657
666
  ncoords[k] = np.array(v)[idx[i]].squeeze()
658
667
 
659
- # revert to standard numpy array if we weren't able to keep coords consistent with the numpy array data
660
- if not check_shapes(obj.shape, ncoords.shape):
668
+ # revert to standard numpy array if we weren't able to keep coords consistent with the numpy array data.
669
+ # This commonly happens with advanced indexing (or pair-wise indexing)
670
+ if not utils.check_shapes(obj.shape, ncoords.shape):
661
671
  return obj.view(np.ndarray)
662
672
 
663
673
  # if dim and the obj shape match, update the dim member of the indexed obj and return
664
674
  obj.coords = ncoords
665
675
  return obj
666
-
676
+
667
677
  # if the coords were unable to be indexed, clear the coords and return a unlabeled numpy array.
668
678
  except Exception:
669
679
  obj.coords = None
@@ -796,9 +806,17 @@ class ldarray(np.ndarray):
796
806
 
797
807
  # convert coordinate to standard index
798
808
  if isinstance(v, (list, tuple, np.ndarray)):
809
+ v = np.atleast_1d(v)
799
810
  # get standard indices for each value in list
800
- np_index[np_i] = [handler(vv, coords_k, **handler_kwargs) for vv in v]
801
-
811
+ np_index[np_i] = np.reshape([handler(vv, coords_k, **handler_kwargs) for vv in v.flatten()], v.shape)
812
+ # recast as labeled array
813
+ if isinstance(v, ldarray):
814
+ np_index[np_i] = ldarray(np_index[np_i], v.coords)
815
+ # cast single valued arrays as slices, this preserves the dimension
816
+ if np_index[np_i].size == 1:
817
+ idx_v = np_index[np_i].item()
818
+ np_index[np_i] = slice(idx_v, idx_v+1)
819
+
802
820
  elif isinstance(v, slice):
803
821
  # call handler for each start, stop and step value
804
822
  s_start, s_stop = [handler(vv, coords_k, **handler_kwargs) if vv is not None else None for vv in [v.start, v.stop]]
@@ -812,27 +830,67 @@ class ldarray(np.ndarray):
812
830
 
813
831
  # if more than one index is a list or array, numpy does pair-wise indexing. Otherwise, we can return the
814
832
  # indices as is.
815
- if np.count_nonzero([isinstance(idx, list) for idx in np_index]) <= 1:
816
- return tuple(np_index)
817
-
818
- # create pairwise indices.
819
- for i, idx in enumerate(np_index):
820
- # convert slice indices to a range of indices
821
- if isinstance(np_index[i], slice):
833
+ # shape of each index
834
+ is_idx_2d = [len(idx.shape) > 1 if isinstance(idx, (np.ndarray)) else False for idx in np_index]
835
+ is_idx_vector = [len(idx) > 1 if isinstance(idx, (list, tuple, np.ndarray)) else False for idx in np_index]
822
836
 
823
- start = 0 if idx.step is None else idx.start
824
- stop = self.shape[i] if idx.stop is None else idx.stop + 1
825
- step = 1 if idx.step is None else idx.step
837
+ if np.any(is_idx_2d):
838
+ # create advanced pairwise indices (or advanced indices). Every index is an array of the same shape,
839
+ # or will be broadcast together so they are the same shape. The indexing arrays must be labeled
840
+ # and have all the dimensions present, minus the dimension it is selecting.
826
841
 
827
- np_index[i] = np.arange(start, stop, step)
842
+ # get the first 2D matrix index
843
+ idx_2d = np_index[is_idx_2d.index(1)]
844
+ result_shape = idx_2d.shape
828
845
 
829
- else:
830
- np_index[i] = np.atleast_1d(idx)
846
+ # check that it is labeled
847
+ if not isinstance(idx_2d, ldarray):
848
+ raise ValueError("Matrix indices must be labeled numpy arrays with all dimensions present.")
849
+
850
+ for i, idx in enumerate(np_index):
851
+
852
+ # create matrix indices for dimensions left with full slices (:)
853
+ if isinstance(np_index[i], slice):
854
+
855
+ key = dim_keys[i]
856
+
857
+ start = 0 if idx.step is None else idx.start
858
+ stop = self.shape[i] if idx.stop is None else idx.stop + 1
859
+ step = 1 if idx.step is None else idx.step
860
+
861
+ idx_b = [None] * len(result_shape)
862
+ # place of index in the result array
863
+ result_place = list(idx_2d.coords.keys()).index(key)
864
+ idx_b[result_place] = slice(None)
865
+ # add extra dimensions
866
+ np_index[i] = np.array(np.arange(start, stop, step))[tuple(idx_b)]
867
+
868
+ # return a meshgrid of index values, the resulting array when this index is used will have the same
869
+ # shape as each array in the axis positions. np.ix_ doesn't perform a full meshgrid broadcast, but ensures
870
+ # the shapes are compatible.
871
+ return tuple([np.broadcast_to(m, result_shape) for m in np_index])
872
+
873
+ # if more than one index is a vector, advanced indexing is used. Broadcast the indices together.
874
+ elif np.count_nonzero(is_idx_vector) > 1:
875
+
876
+ for i, idx in enumerate(np_index):
877
+ # convert slice indices to a range of indices
878
+ if isinstance(np_index[i], slice):
879
+
880
+ start = 0 if idx.step is None else idx.start
881
+ stop = self.shape[i] if idx.stop is None else idx.stop + 1
882
+ step = 1 if idx.step is None else idx.step
883
+
884
+ np_index[i] = np.arange(start, stop, step)
885
+
886
+ else:
887
+ np_index[i] = np.atleast_1d(idx)
888
+
889
+ return np.ix_(*np_index)
890
+
891
+ else:
892
+ return tuple(np_index)
831
893
 
832
- # return a meshgrid of index values, the resulting array when this index is used will have the same
833
- # shape as each array in the axis positions. np.ix_ doesn't perform a full meshgrid broadcast, but ensures
834
- # the shapes are compatible.
835
- return np.ix_(*np_index)
836
894
 
837
895
  def save(self, filepath: str):
838
896
  """
@@ -885,6 +943,8 @@ class ldarray(np.ndarray):
885
943
  cval: float = 0,
886
944
  prefilter: bool = True,
887
945
  dtype: np.dtype = None,
946
+ precision: int = 6,
947
+ flat: bool = False,
888
948
  **coords,
889
949
  ):
890
950
  """
@@ -901,12 +961,16 @@ class ldarray(np.ndarray):
901
961
  The mode parameter determines how the input array is extended beyond its boundaries. Default is "constant".
902
962
  cval : float, default: 0.0
903
963
  Value to fill past edges of input if mode is "constant". Default is 0.0.
904
- prefilter : bool, default: False
964
+ prefilter : bool, default: True
905
965
  Determines if the input array is prefiltered with spline_filter before interpolation.
906
- The default is False.
907
966
  dtype : np.dtype, optional
908
967
  The dtype of the returned array. By default, the dtype is the same as the input array, which may lead to
909
968
  unexpected results if interpolating an integer array.
969
+ precision : int, optional
970
+ decimal precision of interpolation, default is 6 decimal places.
971
+ flat : bool, default: False
972
+ flattens all coords into a pairwise interpolation if True. Default is False, which creates
973
+ a grid interpolation across all coords.
910
974
  **coords
911
975
  coordinate values to interpolate at. Each value is typically a 1D vector of coordinate values, but
912
976
  multi-dimensional arrays are also supported if they are provided as an ldarray. The interpolated
@@ -958,12 +1022,22 @@ class ldarray(np.ndarray):
958
1022
 
959
1023
  """
960
1024
 
1025
+ # set any nan values to 0
1026
+ data = np.nan_to_num(self)
1027
+
961
1028
  coords = {k: np.atleast_1d(v) for k, v in coords.items()}
1029
+ v0 = list(coords.values())[0]
962
1030
 
1031
+ if flat:
1032
+ if not np.all([len(v) == len(v0) for v in coords.values()]):
1033
+ raise ValueError("All pair-wise interpolation coords must be equal length.")
1034
+
963
1035
  # coordinate keys that are specified as meshgrids
964
1036
  mg_keys = [k for k in self.coords.keys() if k in coords.keys() and len(coords[k].shape) > 1]
965
1037
  # dimension indices for all coordinates that are single vectors and not meshgrids
966
1038
  vector_idx = [i for i, k in enumerate(self.coords.keys()) if k not in mg_keys]
1039
+ # dimensions that do not have interp coords
1040
+ missing_idx = [i for i, k in enumerate(self.coords.keys()) if k not in coords.keys()]
967
1041
 
968
1042
  # check that all meshgrid indices have the same shape
969
1043
  if len(mg_keys):
@@ -972,17 +1046,8 @@ class ldarray(np.ndarray):
972
1046
  raise ValueError("All meshgrid indices must be the same shape.")
973
1047
 
974
1048
  # all meshgrids must be labeled with the same coordinates
975
- if not all([isinstance(coords[k], ldarray) and coords[k].coords == m0.coords for k in mg_keys]):
976
- raise ValueError("All meshgrid indices must labeled arrays with identical coordinates.")
977
-
978
- # interpolated shape is the length of each data coordinates that are given as vectors (or not included),
979
- # followed by the meshgrid shape.
980
- dim_keys = list(self.coords.keys())
981
- interp_shape = tuple(
982
- [self.shape[i] if dim_keys[i] not in coords.keys() else len(coords[dim_keys[i]]) for i in vector_idx]
983
- )
984
- if len(mg_keys):
985
- interp_shape += m0.shape
1049
+ if not isinstance(m0, ldarray):
1050
+ raise ValueError("Meshgrid indices must labeled arrays with identical coordinates.")
986
1051
 
987
1052
  # Start with list of slices that index the full range of each dimension.
988
1053
  # dimensions that are not included in coords will be left as a full vector of all indices in
@@ -1014,20 +1079,37 @@ class ldarray(np.ndarray):
1014
1079
 
1015
1080
  # get the floating point "index" by interpolation for each coordinate value.
1016
1081
  else:
1017
- coord_interp = interp1d(coords_k, np.arange(0, self.shape[np_i]), assume_sorted=False, kind="linear")
1018
- interp_index[np_i] = coord_interp(v)
1082
+ coord_interp = interpolate.interp1d(
1083
+ coords_k, np.arange(0, self.shape[np_i]), assume_sorted=False, kind="linear"
1084
+ )
1085
+ interp_index[np_i] = coord_interp(np.around(v, decimals=precision))
1019
1086
 
1020
1087
  # map_coordinates work similarly as numpy advanced indexing, where the index for each dimension can
1021
1088
  # be an matrix. The matrices must all be the same shape, so broadcast the matrices/vectors in interp_index
1022
1089
  # across each other. The number of interpolated dimensions does not need to be the same as the array dimensions.
1023
1090
  interp_index_b = [None] * self.ndim
1024
- v_i = 0
1025
1091
 
1092
+ if flat:
1093
+ # interpolated shape is the shape of the dimensions that are not in the interp coordinates,
1094
+ # plus the flat interpolated vector.
1095
+ interp_shape = [len(v) for k, v in self.coords.items() if k not in coords.keys()] + [len(v0)]
1096
+
1097
+
1098
+ else:
1099
+ # interpolated shape is the length of each data coordinates that are given as vectors (or not included),
1100
+ # followed by the meshgrid shape.
1101
+
1102
+ interp_shape = tuple(
1103
+ [self.shape[i] if dim_keys[i] not in coords.keys() else len(coords[dim_keys[i]]) for i in vector_idx]
1104
+ )
1105
+ if len(mg_keys):
1106
+ interp_shape += m0.shape
1107
+
1108
+ v_i = 0
1026
1109
  for i in range(self.ndim):
1027
1110
 
1028
- # for vector indices, add dimensions for all the other vector dimensions, as well as the meshgrid
1029
- # dimensions.
1030
- if i in vector_idx:
1111
+ # for vector indices, add dimensions for all the other dimensions
1112
+ if (flat and i in missing_idx) or (not flat and i in vector_idx):
1031
1113
  # select current dimension in the interpolated shape by adding a ":" in the dimension list.
1032
1114
  # the vector indices are stacked at the front of the interpolated shape, regardless of where
1033
1115
  # they appear in the array dimensions (use v_i instead of i to select dimension)
@@ -1036,9 +1118,12 @@ class ldarray(np.ndarray):
1036
1118
  # add extra dimensions
1037
1119
  interp_index_b[i] = np.array(interp_index[i])[tuple(idx_b)]
1038
1120
  v_i += 1
1039
- # for meshgrid indices, add extra dimensions for the vector dimensions at the beginning of the array
1121
+
1122
+ # for meshgrid indices or flattened dimensions, add extra dimensions for the vector dimensions
1123
+ # at the beginning of the array.
1040
1124
  else:
1041
- interp_index_b[i] = interp_index[i][tuple([None] * v_i)]
1125
+ interp_index_b[i] = interp_index[i][tuple([None] * v_i)] if v_i else interp_index[i]
1126
+
1042
1127
 
1043
1128
  # map_coordinates doesn't broadcast the indices like numpy does for advanced indexing. Broadcast
1044
1129
  # index array to the same shape for each dimension.
@@ -1047,24 +1132,130 @@ class ldarray(np.ndarray):
1047
1132
  if dtype is None:
1048
1133
  dtype = self.dtype
1049
1134
 
1050
- data = ndimage.map_coordinates(
1051
- self.astype(dtype), map_idx, output=output, order=order, mode=mode, cval=cval, prefilter=prefilter
1135
+ data_interp = ndimage.map_coordinates(
1136
+ data.astype(dtype), map_idx, output=output, order=order, mode=mode, cval=cval, prefilter=prefilter
1052
1137
  )
1053
1138
 
1054
1139
  data_coords = {}
1140
+ attrs = dict()
1055
1141
  # add coordinates from vector indices
1056
1142
  for i, k in enumerate(self.coords.keys()):
1057
- if i in vector_idx:
1143
+ if (flat and i in missing_idx) or (not flat and i in vector_idx):
1058
1144
  data_coords[k] = coords[k] if k in coords.keys() else self.coords[k]
1059
1145
 
1060
1146
  # add the coordinates from the meshgrid
1061
1147
  if len(mg_keys):
1062
1148
  data_coords.update(m0.coords)
1063
1149
 
1150
+ # add the coordinates for the flattened dimensions
1151
+ if flat:
1152
+ flat_key = "".join(coords.keys())
1153
+ data_coords[flat_key] = np.arange(len(v0))
1154
+ attrs = {k: v for k, v in coords.items()}
1155
+
1064
1156
  return ldarray(
1065
- data, coords=data_coords
1157
+ data_interp, coords=data_coords, attrs=attrs
1066
1158
  )
1067
1159
 
1160
+
1161
+ def interpolate_from_flat(self, flat: bool = False, **coords):
1162
+ """
1163
+ Interpolate pairwise, flattened dimensions. Two (and only two) interpolation dimensions are supported.
1164
+ The pairwise coordinates must be present in the attributes.
1165
+
1166
+ Parameters
1167
+ ----------
1168
+ **coords
1169
+ coordinate values to interpolate at. If dimension is "uv", interpolated coords must be
1170
+ "u" and "v".
1171
+
1172
+ flat: bool, default: False
1173
+ if False (default), the data is returned as a meshgrid of the two interpolation coordinates.
1174
+ If True, the data is returned as pairwise points of the interpolation coordinate.
1175
+
1176
+ Examples
1177
+ --------
1178
+ """
1179
+
1180
+ interp_keys = list(coords.keys())
1181
+ interp_v1, interp_v2 = [np.atleast_1d(v) for v in coords.values()]
1182
+
1183
+ if len(interp_keys) != 2:
1184
+ raise ValueError("Requires a pair of interpolation coordinates.")
1185
+
1186
+ # ensure shape of coords matches if flat
1187
+ if (len(interp_v1.shape) > 1) and (interp_v1.shape != interp_v2.shape):
1188
+ raise ValueError("Interpolation coordinates must have equal shapes.")
1189
+
1190
+ # coords must be labeled if more than 1D
1191
+ if len(interp_v1.shape) > 1 and not isinstance(interp_v1, ldarray):
1192
+ raise ValueError("Interpolation coordinates must be labeled numpy arrays if greater than 1D.")
1193
+
1194
+ if flat:
1195
+ if not len(interp_v1) == len(interp_v1):
1196
+ raise ValueError("All pair-wise interpolation coords must be equal length.")
1197
+
1198
+ # get data coordinates for both interpolated dimensions
1199
+ if all([k in self.coords.keys() for k in interp_keys]):
1200
+ data = self.transpose((*interp_keys, ...))
1201
+ data_coords_m = np.meshgrid(*[self.coords[k] for k in interp_keys], indexing="ij")
1202
+ # flatten and stack mesh so coords are Nx2
1203
+ data_coords = np.stack(data_coords_m, axis=-1).reshape((-1, 2))
1204
+ # flatten interpolated coords in data
1205
+ data = np.reshape(data, (len(data_coords), *data.shape[2:]))
1206
+
1207
+ # if data coordinates are flattened into one dimensions, use the attributes to create Nx2 positions
1208
+ elif "".join(interp_keys) in self.coords.keys():
1209
+ data = self.transpose(("".join(interp_keys), ...))
1210
+ data_coords = np.stack([self.attrs[k] for k in interp_keys], axis=-1)
1211
+
1212
+ else:
1213
+ raise ValueError(f"Unable to interpolate coordinates {list(self.coords.keys())}")
1214
+
1215
+ # create interpolator, this does handle complex data but performs better if interpolation is done
1216
+ # on magnitude and angle separately.
1217
+ # set any nan values to 0
1218
+ data = np.nan_to_num(data)
1219
+ # leave extrapolated values at nan
1220
+ abs_data = np.abs(data)
1221
+ interp_func_mag = interpolate.CloughTocher2DInterpolator(data_coords, abs_data, fill_value=np.nan)
1222
+ interp_func_phasor = interpolate.CloughTocher2DInterpolator(data_coords, data / abs_data, fill_value=np.nan)
1223
+
1224
+ # stack coordinates so shape is ..., 2
1225
+ if len(interp_v1.shape) > 1 or flat:
1226
+ interp_pos = np.stack((interp_v1, interp_v2), axis=-1).reshape((-1, 2))
1227
+ else:
1228
+ interp_pos_m = np.meshgrid(interp_v1, interp_v2, indexing="ij")
1229
+ interp_pos = np.stack(interp_pos_m, axis=-1).reshape((-1, 2))
1230
+
1231
+ # create result coords
1232
+ if flat:
1233
+ interp_coords = {"".join(coords.keys()): np.arange(0, len(interp_v1))}
1234
+ # interpolated coords are the same as the argument coords if a meshgrid was passed in
1235
+ elif len(interp_v1.shape) > 1:
1236
+ interp_coords = interp_v1.coords
1237
+ # otherwise use the argument values as coords
1238
+ else:
1239
+ interp_coords = coords
1240
+
1241
+ # evaluate interpolation
1242
+ with np.errstate(all="ignore"):
1243
+ phasor = interp_func_phasor(interp_pos)
1244
+ data_interp = interp_func_mag(interp_pos) * (phasor / np.abs(phasor))
1245
+
1246
+ interp_data = ldarray(
1247
+ data_interp.reshape(*[len(v) for v in interp_coords.values()], *data.shape[1:]),
1248
+ coords = dict(
1249
+ **interp_coords, **{k: v for k, v in self.coords.items() if k not in (*interp_keys, "".join(interp_keys))}
1250
+ )
1251
+ )
1252
+
1253
+ # add flattened coordinates as attributes
1254
+ if flat:
1255
+ interp_data.attrs = {k: v for k, v in coords.items()}
1256
+
1257
+ return interp_data
1258
+
1068
1259
  @classmethod
1069
1260
  def load(cls, filepath: str, **kwargs):
1070
1261
  """
@@ -1148,6 +1339,260 @@ class ldarray(np.ndarray):
1148
1339
 
1149
1340
  return ldarray(super().transpose(order_idx), coords=coords)
1150
1341
 
1342
+ def plot(
1343
+ self,
1344
+ xaxis: str = None,
1345
+ xfmt: str = "real",
1346
+ yfmt: str = "real",
1347
+ legend: bool = True,
1348
+ ax = None,
1349
+ lines = None,
1350
+ ymin: float = None,
1351
+ ymax: float = None,
1352
+ format_axes: bool = True,
1353
+ **kwargs
1354
+ ):
1355
+ """
1356
+ Create line plot for labeled numpy array.
1357
+
1358
+ Parameters
1359
+ ----------
1360
+ xaxis : str, optional
1361
+ dimension to plot along the x-axis, chooses the first dimension if not provided.
1362
+
1363
+ ax : plt.Axes, optional
1364
+ matplotlib axes object
1365
+
1366
+ xfmt : (np.ndarray) -> np.ndarray, optional
1367
+ String value that determines how to format the x-axis data before plotting.
1368
+ An arbitrary function is also supported that accepts a 1D numpy array and returns a formatted array.
1369
+
1370
+ The following string values are supported for the xmft or yfmt arguments:
1371
+ - "db20" : `20 * np.log10(...)`
1372
+ - "db10" : `10 * np.log10(...)`
1373
+ - "abs" : `np.abs(...)`
1374
+ - "mag" : `np.abs(...)`
1375
+ - "deg" : `np.angle(..., deg=True)`
1376
+ - "rad" : `np.angle(..., deg=False)`
1377
+ - "angle": `np.angle(..., deg=False)`
1378
+ - "real" : `np.real(...)`
1379
+ - "imag" : `np.imag(...)`
1380
+ - "deg2rad" : `np.deg2rad(...)`
1381
+ - "rad2deg" : `np.rad2deg(...)`
1382
+
1383
+ yfmt : (np.ndarray) -> np.ndarray, optional
1384
+ String value that determines how to format the y-axis data before plotting.
1385
+ An arbitrary function is also supported that accepts a 1D numpy array and returns a formatted array.
1386
+
1387
+ **kwargs
1388
+ keys that are in coordinates are passed to .sel(). Remaining kwargs are passed to ax.plot()
1389
+
1390
+ """
1391
+
1392
+ # create axes if one is not provided
1393
+ if ax is None:
1394
+ import matplotlib.pyplot as plt
1395
+ ax = plt.gca()
1396
+
1397
+ # plot along first dimension by default
1398
+ if xaxis is None:
1399
+ xaxis = list(self.coords.keys())[0]
1400
+
1401
+ # select a format function from one of the defaults if provided as a string
1402
+ ylabel = ""
1403
+ yfmt_str = yfmt
1404
+ if isinstance(yfmt, str):
1405
+ ylabel = yfmt
1406
+ yfmt = utils.DATA_FMT_FUNC[yfmt]
1407
+
1408
+ if isinstance(xfmt, str):
1409
+ xfmt = utils.DATA_FMT_FUNC[xfmt]
1410
+
1411
+ # xaxis coords
1412
+ xaxis_coords = xfmt(self.coords[xaxis])
1413
+
1414
+ # select data
1415
+ sel_coords = {k: np.atleast_1d(kwargs.pop(k)) for k in self.coords.keys() if k in kwargs.keys()}
1416
+ data = self.sel(**sel_coords)
1417
+
1418
+ # coords with more than one value (other than the x-axis)
1419
+ other_coords = {k: v for k, v in data.coords.items() if k != xaxis and len(v) > 1}
1420
+ # coords with only one value, these will not be included in legend since they're the same for all lines
1421
+ unitary_coords = {k: v for k, v in data.coords.items() if k != xaxis and len(v) == 1}
1422
+ # label for title with all unitary coords
1423
+ unitary_label = ", ".join([utils.format_label(k, v.item()) for k, v in unitary_coords.items()])
1424
+
1425
+ # all combinations of coordinates
1426
+ combinations = list(product(*other_coords.values()))
1427
+
1428
+ lines_new = []
1429
+
1430
+ for i, comb_i in enumerate(combinations):
1431
+ # get single combination, only dimension should be x-axis
1432
+ coords_dict = {k: comb_i[i] for (i, k) in enumerate(other_coords.keys())}
1433
+ ln_data = data.sel(**coords_dict).squeeze()
1434
+
1435
+ # update line data if lines were provided
1436
+ if lines is not None:
1437
+ lines[i].set_ydata(yfmt(ln_data))
1438
+ # add lines to plot
1439
+ else:
1440
+ # build legend label
1441
+ label = ", ".join([utils.format_label(k, v) for k, v in coords_dict.items()])
1442
+ lines_new += ax.plot(xaxis_coords, yfmt(ln_data), label=label, **kwargs)
1443
+
1444
+ if lines is None and format_axes:
1445
+
1446
+ if legend and len(combinations) > 1 and len(combinations) < 7:
1447
+ ax.legend()
1448
+
1449
+ ax.set_xlabel(xaxis)
1450
+ ax.set_title(f"{unitary_label}", fontsize="medium")
1451
+ ax.grid(True)
1452
+
1453
+ # set upper/lower limit to a multiple of 5 for dB plot.
1454
+ if yfmt_str in ("db20", "db10"):
1455
+ if ymax is None:
1456
+ ymax = np.ceil(np.nanmax(yfmt(data)) / 5) * 5
1457
+ if ymin is None:
1458
+ ymin = np.floor(np.nanmin(yfmt(data)) / 5) * 5
1459
+ # clip to - 40dB range
1460
+ ymin = np.clip(ymin, ymax - 40, None)
1461
+
1462
+ ymin = ax.get_ylim()[0] if ymin is None else ymin
1463
+ ymax = ax.get_ylim()[1] if ymax is None else ymax
1464
+
1465
+ try:
1466
+ ax.set_ylim((ymin, ymax))
1467
+ except:
1468
+ pass
1151
1469
 
1152
-
1470
+ # if polar axes, add the ylabel to the last tick marker
1471
+ if ax.name == "polar":
1472
+ ax.set_theta_zero_location('N')
1473
+ ax.set_theta_direction(-1)
1474
+
1475
+ # polar always interprets the data in radians, set the plot range to be from -180° to 180°.
1476
+ ax.set_thetalim(-np.pi, np.pi)
1477
+ ax.set_thetagrids(range(-180, 180, 45))
1478
+ ax.tick_params(labelsize='small')
1479
+
1480
+ # add label to last tick marker
1481
+ labels = [f"{t:.0f}" for t in ax.get_yticks()]
1482
+ labels[-1] += ("dB" if yfmt_str in ("db20", "db10") else ylabel[:3])
1483
+ ax.set_yticks(ax.get_yticks(), labels)
1484
+
1485
+ # setup cartesian axes limits
1486
+ else:
1487
+ ax.set_ylabel(ylabel)
1488
+ ax.set_xlim([np.nanmin(xaxis_coords), np.nanmax(xaxis_coords)])
1489
+
1490
+ return lines_new
1491
+ else:
1492
+ return lines
1493
+
1494
+ def pcolormesh(
1495
+ self,
1496
+ xaxis: str,
1497
+ yaxis: str,
1498
+ xfmt: str = "real",
1499
+ yfmt: str = "real",
1500
+ zfmt: str = "real",
1501
+ ax = None,
1502
+ mesh = None,
1503
+ colorbar : dict = True,
1504
+ **kwargs
1505
+ ):
1506
+ """
1507
+ Create pcolormesh plot for labeled numpy array.
1508
+
1509
+ Parameters
1510
+ ----------
1511
+ xaxis : str, optional
1512
+ dimension to plot along the x-axis
1513
+
1514
+ yaxis : str, optional
1515
+ dimension to plot along the y-axis
1516
+
1517
+ xfmt : (np.ndarray) -> np.ndarray, optional
1518
+ String value that determines how to format the x-axis data before plotting.
1519
+ An arbitrary function is also supported that accepts a 1D numpy array and returns a formatted array.
1520
+
1521
+ The following string values are supported for the xfmt or yfmt arguments:
1522
+ - "db20" : `20 * np.log10(...)`
1523
+ - "db10" : `10 * np.log10(...)`
1524
+ - "abs" : `np.abs(...)`
1525
+ - "deg" : `np.angle(..., deg=True)`
1526
+ - "rad" : `np.angle(..., deg=False)`
1527
+ - "angle": `np.angle(..., deg=False)`
1528
+ - "real" : `np.real(...)`
1529
+ - "imag" : `np.imag(...)`
1530
+ - "deg2rad" : `np.deg2rad(...)`
1531
+ - "rad2deg" : `np.rad2deg(...)`
1532
+
1533
+ yfmt : (np.ndarray) -> np.ndarray, optional
1534
+ String value that determines how to format the y-axis data before plotting.
1535
+ An arbitrary function is also supported that accepts a 1D numpy array and returns a formatted array.
1536
+
1537
+ zfmt : (np.ndarray) -> np.ndarray, optional
1538
+ String value that determines how to format the z-axis data before plotting.
1539
+ An arbitrary function is also supported that accepts a 2D numpy array and returns a formatted array.
1540
+
1541
+
1542
+ **kwargs
1543
+ keys that are in coordinates are passed to .sel(). Remaining kwargs are passed to ax.pcolormesh()
1544
+
1545
+ """
1546
+
1547
+ # create axes if one is not provided
1548
+ if ax is None:
1549
+ import matplotlib.pyplot as plt
1550
+ ax = plt.gca()
1551
+
1552
+ if colorbar is True:
1553
+ colorbar = dict()
1554
+
1555
+ # select a format function from one of the defaults if provided as a string
1556
+ zlabel = ""
1557
+ if isinstance(zfmt, str):
1558
+ zlabel = zfmt
1559
+ zfmt = utils.DATA_FMT_FUNC[zfmt]
1560
+
1561
+ if isinstance(xfmt, str):
1562
+ xfmt = utils.DATA_FMT_FUNC[xfmt]
1563
+
1564
+ if isinstance(yfmt, str):
1565
+ yfmt = utils.DATA_FMT_FUNC[yfmt]
1566
+
1567
+ # select data
1568
+ sel_coords = {k: np.atleast_1d(kwargs.pop(k)) for k in self.coords.keys() if k in kwargs.keys()}
1569
+ data = self.sel(**sel_coords)
1570
+
1571
+ # data must have only 2 dimensions at this point
1572
+ if data.squeeze().ndim != 2:
1573
+ raise ValueError("Data must have only 2 dimensions.")
1574
+
1575
+ # coords with only one value, these will not be included in legend since they're the same for all lines
1576
+ unitary_coords = {k: v for k, v in data.coords.items() if k != xaxis and len(v) == 1}
1577
+ # label for title with all unitary coords
1578
+ unitary_label = ", ".join([utils.format_label(k, v.item()) for k, v in unitary_coords.items()])
1579
+
1580
+ data = data.squeeze().transpose((yaxis, xaxis))
1581
+
1582
+ # add new colormesh object to plot
1583
+ if mesh is None:
1584
+ mesh = ax.pcolormesh(xfmt(data.coords[xaxis]), yfmt(data.coords[yaxis]), zfmt(data), **kwargs)
1585
+
1586
+ if isinstance(colorbar, dict):
1587
+ c_kwargs = dict(label=zlabel, **colorbar) if "label" not in colorbar.keys() else colorbar
1588
+ ax.figure.colorbar(mesh, **c_kwargs)
1589
+
1590
+ # update existing colormesh
1591
+ else:
1592
+ mesh.set_array(zfmt(data))
1593
+
1594
+ ax.set_xlabel(xaxis)
1595
+ ax.set_title(f"{unitary_label}", fontsize="medium")
1596
+ ax.set_ylabel(yaxis)
1153
1597
 
1598
+ return mesh
@@ -0,0 +1,74 @@
1
+ import numpy as np
2
+
3
+ DATA_FMT_FUNC = {
4
+ "mag": np.abs,
5
+ "abs": np.abs,
6
+ "db20": lambda x: 20 * np.log10(np.abs(x)),
7
+ "db10": lambda x: 10 * np.log10(np.abs(x)),
8
+ "deg": lambda x: np.angle(x, deg=True),
9
+ "rad": lambda x: np.angle(x, deg=False),
10
+ "angle": lambda x: np.angle(x, deg=False),
11
+ "deg_unwrap": lambda x: np.rad2deg(np.unwrap(np.angle(x))),
12
+ "real": np.real,
13
+ "imag": np.imag,
14
+ "deg2rad" : np.deg2rad,
15
+ "rad2deg" : np.rad2deg
16
+ }
17
+
18
+ LABEL_FMT_FUNC = dict(
19
+ frequency = lambda x: f"{x/1e9:.3f}GHz"
20
+ )
21
+
22
+
23
+ def check_shapes(a: tuple, b: tuple):
24
+ """
25
+ Check that the shape tuples a and b match
26
+ """
27
+
28
+ if len(a) != len(b):
29
+ return False
30
+
31
+ # check that the length of each dimension matches
32
+ return all([a[i] == b[i] for i in range(len(a))])
33
+
34
+ def check_coords(c1: dict, c2: dict, tolerance=1e-6):
35
+ """
36
+ Check that two coordinates are identical
37
+ """
38
+ assert tuple(c1.keys()) == tuple(c2.keys()), f"Coord keys are different: {c1.keys()} vs {c2.keys()}"
39
+
40
+ for k in c1.keys():
41
+ if np.any(np.abs(c1[k] - c2[k]) > tolerance):
42
+ return False
43
+
44
+ return True
45
+
46
+ def add_data_formatter(key: str, func):
47
+ """
48
+ Add a custom data format function. Function must accept a single numpy array argument and return a
49
+ formatted array of the same shape.
50
+ """
51
+ DATA_FMT_FUNC[key] = func
52
+
53
+ def add_label_formatter(key: str, func):
54
+ """
55
+ Add a custom data format function. Function must accept a single scalar argument and return a
56
+ formatted string
57
+ """
58
+ LABEL_FMT_FUNC[key] = func
59
+
60
+ def format_label(
61
+ key: str, value
62
+ ) -> str:
63
+
64
+ if key in LABEL_FMT_FUNC.keys():
65
+ return LABEL_FMT_FUNC[key](value)
66
+
67
+ # create default label formatters if not included in look up table
68
+ if isinstance(value, (float, np.floating)):
69
+ return f"{key}={value:.3f}"
70
+ elif isinstance(value, (int, np.integer)):
71
+ return f"{key}={value}"
72
+ # don't include key in label for string coordinates
73
+ else:
74
+ return "{}".format(value)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: np-struct
3
- Version: 0.1.6
3
+ Version: 0.2.0
4
4
  Summary: Interface for NumPy structured arrays
5
5
  Author-email: Rick Lyon <rlyon14@gmail.com>
6
6
  Project-URL: repository, https://github.com/ricklyon/np_struct
@@ -6,6 +6,7 @@ np_struct/bitfields.py
6
6
  np_struct/ldarray.py
7
7
  np_struct/structures.py
8
8
  np_struct/transfer.py
9
+ np_struct/utils.py
9
10
  np_struct.egg-info/PKG-INFO
10
11
  np_struct.egg-info/SOURCES.txt
11
12
  np_struct.egg-info/dependency_links.txt
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "np-struct"
7
- version = "0.1.6"
7
+ version = "0.2.0"
8
8
  description = "Interface for NumPy structured arrays"
9
9
  readme = "README.md"
10
10
  authors = [{ name = "Rick Lyon", email = "rlyon14@gmail.com" }]
@@ -16,7 +16,7 @@ classifiers = [
16
16
  ]
17
17
  keywords = ["numpy", "structured array", "struct"]
18
18
  dependencies = [
19
- "numpy", "scipy"
19
+ "numpy", "scipy",
20
20
  ]
21
21
  requires-python = ">=3.7"
22
22
 
@@ -30,7 +30,7 @@ include = ["np_struct"]
30
30
  repository = "https://github.com/ricklyon/np_struct"
31
31
 
32
32
  [tool.bumpversion]
33
- current_version = "0.1.6"
33
+ current_version = "0.2.0"
34
34
  commit = true
35
35
  tag = true
36
36
 
@@ -1,11 +1,14 @@
1
1
  import unittest
2
- from np_struct import ldarray, Coords
2
+ from np_struct import ldarray, Coords, utils
3
3
  import numpy as np
4
4
  from numpy import testing as npt
5
5
  import datetime as dt
6
6
  from dateutil import relativedelta as rdt
7
7
  import os
8
8
 
9
+ from scipy import ndimage
10
+ import numpy as np
11
+
9
12
 
10
13
  class TestLdArray(unittest.TestCase):
11
14
 
@@ -39,6 +42,7 @@ class TestLdArray(unittest.TestCase):
39
42
  npt.assert_array_equal(ld_2, data[np.array([1, 0]), 0])
40
43
  npt.assert_array_equal(ld_2.coords["a"], [1, 0])
41
44
 
45
+
42
46
  def test_float_index(self):
43
47
 
44
48
  coords = dict(a=['data1', 'data2'], b=np.arange(0, 20, 0.2))
@@ -146,7 +150,7 @@ class TestLdArray(unittest.TestCase):
146
150
  np.testing.assert_array_almost_equal(ld_int.sel(a="exp(t)"), np.exp(1j * t_int), decimal=2)
147
151
  np.testing.assert_array_almost_equal(ld_int.sel(a="exp(0.5t)"), np.exp(0.5 * 1j * t_int), decimal=2)
148
152
 
149
- def test_interpolation_2d(self):
153
+ def test_ldarray_indexing(self):
150
154
  coords = dict(a=[1, 2], b=['data1', 'data2', 'data3'])
151
155
  ld = ldarray([[10, 8, 6], [0, 2, 4]], coords=coords)
152
156
 
@@ -167,6 +171,105 @@ class TestLdArray(unittest.TestCase):
167
171
  np.testing.assert_array_equal(data.coords["x"], [0, 1])
168
172
  np.testing.assert_array_equal(data.coords["y"], [0, 1])
169
173
 
174
+ def test_ldarray_indexing2(self):
175
+ # test index values that are ldarrays
176
+
177
+ coords = dict(a=np.arange(3), b=np.arange(10, 16), c=np.arange(3))
178
+ data = np.arange(54).reshape(3, 6, 3)
179
+ ld = ldarray(data, coords=coords)
180
+
181
+ # index must include all coordinates and have the same shape as the data,
182
+ # except the indexing dimension (in this case "b")
183
+ idx_v = np.ones((3, 2, 2, 3)) * 10
184
+ idx_v[0, 0, :] = 12
185
+ idx_v[2, 1, :] = 11
186
+
187
+ idx = ldarray(idx_v, coords=dict(a=np.arange(3), new1=[12, 11], new2=[6, 7], c=np.arange(3)))
188
+ result = ld.sel(b=idx)
189
+
190
+ # both values of new2 axis should be at b=12 for the a=0, new1=12 dimension
191
+ np.testing.assert_array_almost_equal(result.sel(a=0, new1=12)[0], ld.sel(a=0, b=12))
192
+ np.testing.assert_array_almost_equal(result.sel(a=0, new1=12)[1], ld.sel(a=0, b=12))
193
+ np.testing.assert_array_almost_equal(result.sel(a=0, new1=11)[0], ld.sel(a=0, b=10))
194
+ np.testing.assert_array_almost_equal(result.sel(a=0, new1=11)[1], ld.sel(a=0, b=10))
195
+
196
+ # both new1 and new2 should be at b=10 for the a=1 dimension
197
+ np.testing.assert_array_almost_equal(result.sel(a=1)[0, 0], ld.sel(a=1, b=10))
198
+ np.testing.assert_array_almost_equal(result.sel(a=1)[1, 0], ld.sel(a=1, b=10))
199
+ np.testing.assert_array_almost_equal(result.sel(a=1)[0, 1], ld.sel(a=1, b=10))
200
+ np.testing.assert_array_almost_equal(result.sel(a=1)[1, 1], ld.sel(a=1, b=10))
201
+
202
+ # both values of new 2 should be at b=11 for the a=2, new1=11 dimension
203
+ np.testing.assert_array_almost_equal(result.sel(a=2, new1=12)[0], ld.sel(a=2, b=10))
204
+ np.testing.assert_array_almost_equal(result.sel(a=2, new1=12)[1], ld.sel(a=2, b=10))
205
+ np.testing.assert_array_almost_equal(result.sel(a=2, new1=11)[0], ld.sel(a=2, b=11))
206
+ np.testing.assert_array_almost_equal(result.sel(a=2, new1=11)[1], ld.sel(a=2, b=11))
207
+
208
+ self.assertTrue(utils.check_coords(result.coords, idx.coords), "coords are not equal.")
209
+
210
+ def test_interpolation_nan(self):
211
+ t = np.linspace(0, 2 * np.pi, 21)
212
+ t_int = np.linspace(1, 2, 61)
213
+
214
+ data = np.array([np.exp(1j * t), np.exp(0.5 * 1j * t)])
215
+
216
+ ld = ldarray(data, coords = dict(a=["exp(t)", "exp(0.5t)"], t=t))
217
+
218
+ # set end points to nan
219
+ ld[:, -5:] = np.nan
220
+
221
+ data = ld.interpolate(t=t_int).sel(a="exp(t)")
222
+ np.testing.assert_array_almost_equal(data, np.exp(1j * t_int), decimal=2)
223
+
224
+ def test_interpolation_round_single(self):
225
+
226
+ t = np.linspace(0, 2 * np.pi, 21)
227
+ t_int = np.linspace(0.5, 5.5, 61)
228
+
229
+ data = np.array([np.sin(t)])
230
+ ld = ldarray(data, coords = dict(a=0, t=t))
231
+
232
+ np.testing.assert_array_almost_equal(ld.interpolate(a=[-1e-7], t=t_int)[0], np.sin(t_int), decimal=2)
233
+
234
+ def test_interpolation_flat(self):
235
+ # avoid interpolating at the endpoints, it's close to the right value but hard to test exactly
236
+ t = np.linspace(0, 2 * np.pi, 21)
237
+
238
+ data = np.array([np.sin(t), np.cos(t)])
239
+ ld = ldarray(data, coords = dict(a=["sin", "cos"], t=t))
240
+
241
+ coords = dict(t = [0, 3, 6], a=["sin", "sin", "cos"])
242
+ ld.interpolate(**coords, flat=True)
243
+
244
+ def test_interpolate_2d(self):
245
+ # test 2d interpolator on flattened data
246
+ u = np.linspace(-1, 1, 21)
247
+ v = np.linspace(-1, 1, 21)
248
+
249
+ u_m, v_m = np.meshgrid(u, v)
250
+ z = np.cos(u_m) * np.cos(v_m)
251
+
252
+ data = ldarray(z, coords=dict(u=u, v=v))
253
+
254
+ u_int = np.linspace(-1, 1, 141)
255
+ v_int = np.linspace(-1, 1, 141)
256
+
257
+ # interpolate_from_flat should work on meshgrid data as well. It is slower and not as accurate, but
258
+ # should be reasonably close
259
+ d1 = data.interpolate(u=u_int, v=v_int)
260
+ d2 = data.interpolate_from_flat(u=u_int, v=v_int)
261
+ np.testing.assert_array_less(np.mean(np.abs(d1 - d2)), 0.002)
262
+
263
+ # interpolate with pairwise points
264
+ data_flat = data.interpolate(u=u_m.flatten(), v=v_m.flatten(), flat=True)
265
+
266
+ np.testing.assert_array_almost_equal(np.reshape(data_flat, (21, 21)), data)
267
+ # flattened coords should be labeled 0-N, where N is the number of pairwise points
268
+ np.testing.assert_array_almost_equal(data_flat.uv, np.arange(21*21))
269
+
270
+ # interpolate the flattened data to get back to the original data
271
+ data_round_trip = data_flat.interpolate_from_flat(u=u, v=v)
272
+ np.testing.assert_array_almost_equal(data_round_trip, data)
170
273
 
171
274
  def test_save(self):
172
275
 
File without changes
File without changes
File without changes