@@ -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
9696def 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 )
0 commit comments