Skip to content

Commit d1ee511

Browse files
Reshaping return array of _get_corner_data_Agrid to (ti, zi, yi, xi, npart)
For easier use in other Interpolators
1 parent 741698c commit d1ee511

2 files changed

Lines changed: 25 additions & 16 deletions

File tree

src/parcels/interpolators.py

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -74,11 +74,11 @@ def _get_corner_data_Agrid(
7474

7575
# Y coordinates: [yi, yi, yi+1, yi+1] for each spatial point, repeated for time/z
7676
yi_1 = np.clip(yi + 1, 0, data.shape[2] - 1)
77-
yi = np.tile(np.repeat(np.column_stack([yi, yi_1]), 2), (lenT) * (lenZ))
77+
yi = np.tile(np.array([yi, yi, yi_1, yi_1]).flatten(), lenT * lenZ)
7878

7979
# X coordinates: [xi, xi+1, xi, xi+1] for each spatial point, repeated for time/z
8080
xi_1 = np.clip(xi + 1, 0, data.shape[3] - 1)
81-
xi = np.tile(np.column_stack([xi, xi_1, xi, xi_1]).flatten(), (lenT) * (lenZ))
81+
xi = np.tile(np.array([xi, xi_1]).flatten(), lenT * lenZ * 2)
8282

8383
# Create DataArrays for indexing
8484
selection_dict = {
@@ -90,7 +90,7 @@ def _get_corner_data_Agrid(
9090
if "time" in data.dims:
9191
selection_dict["time"] = xr.DataArray(ti, dims=("points"))
9292

93-
return data.isel(selection_dict).data.reshape(lenT, lenZ, npart, 4)
93+
return data.isel(selection_dict).data.reshape(lenT, lenZ, 2, 2, npart)
9494

9595

9696
def XLinear(
@@ -113,22 +113,22 @@ def XLinear(
113113
corner_data = _get_corner_data_Agrid(data, ti, zi, yi, xi, lenT, lenZ, len(xsi), axis_dim)
114114

115115
if lenT == 2:
116-
tau = tau[np.newaxis, :, np.newaxis]
117-
corner_data = corner_data[0, :, :, :] * (1 - tau) + corner_data[1, :, :, :] * tau
116+
tau = tau[np.newaxis, :]
117+
corner_data = corner_data[0, :] * (1 - tau) + corner_data[1, :] * tau
118118
else:
119-
corner_data = corner_data[0, :, :, :]
119+
corner_data = corner_data[0, :]
120120

121121
if lenZ == 2:
122-
zeta = zeta[:, np.newaxis]
123-
corner_data = corner_data[0, :, :] * (1 - zeta) + corner_data[1, :, :] * zeta
122+
zeta = zeta[np.newaxis, :]
123+
corner_data = corner_data[0, :] * (1 - zeta) + corner_data[1, :] * zeta
124124
else:
125-
corner_data = corner_data[0, :, :]
125+
corner_data = corner_data[0, :]
126126

127127
value = (
128-
(1 - xsi) * (1 - eta) * corner_data[:, 0]
129-
+ xsi * (1 - eta) * corner_data[:, 1]
130-
+ (1 - xsi) * eta * corner_data[:, 2]
131-
+ xsi * eta * corner_data[:, 3]
128+
(1 - xsi) * (1 - eta) * corner_data[0, 0, :]
129+
+ xsi * (1 - eta) * corner_data[0, 1, :]
130+
+ (1 - xsi) * eta * corner_data[1, 0, :]
131+
+ xsi * eta * corner_data[1, 1, :]
132132
)
133133
return value.compute() if is_dask_collection(value) else value
134134

@@ -408,8 +408,8 @@ def _Spatialslip(
408408
corner_dataV = _get_corner_data_Agrid(vectorfield.V.data, ti, zi, yi, xi, lenT, lenZ, npart, axis_dim)
409409

410410
def is_land(ti: int, zi: int, yi: int, xi: int):
411-
uval = corner_dataU[ti, zi, :, xi + 2 * yi]
412-
vval = corner_dataV[ti, zi, :, xi + 2 * yi]
411+
uval = corner_dataU[ti, zi, yi, xi, :]
412+
vval = corner_dataV[ti, zi, yi, xi, :]
413413
return np.where(np.isclose(uval, 0.0) & np.isclose(vval, 0.0), True, False)
414414

415415
f_u = np.ones_like(xsi)

tests/test_interpolation.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -68,9 +68,18 @@ def field():
6868
[0.49, 0.49],
6969
[0.51, 0.51],
7070
[1.49, 6.49],
71-
id="Linear",
71+
id="Linear-1",
7272
),
7373
pytest.param(XLinear, np.timedelta64(1, "s"), 2.5, 0.49, 0.51, 13.99, id="Linear-2"),
74+
pytest.param(
75+
XLinear,
76+
[np.timedelta64(0, "s"), np.timedelta64(1, "s"), np.timedelta64(1, "s")],
77+
[0, 0, 2.5],
78+
[0.49, 0.49, 0.49],
79+
[0.51, 0.51, 0.51],
80+
[1.49, 6.49, 13.99],
81+
id="Linear-3",
82+
),
7483
pytest.param(
7584
XNearest,
7685
[np.timedelta64(0, "s"), np.timedelta64(3, "s")],

0 commit comments

Comments
 (0)