四维数组时间序列中0值最后出现位置的高效计算与代码优化
I have a 4D array with dimensions
[time, model number, longitude, latitude], containing only 0s and 1s. For each time series corresponding to[model, longitude, latitude], I need to compute the last year a 0 appeared, following these rules:
- Return 0 if the entire time series is 0s
- Return 1920 if the entire time series is 1s
- Return the year of the last 0 occurrence only if both 0s and 1s are present
My current nested loop implementation is too slow, and I'm looking for a vectorized, more efficient approach.
Original code:
element=0 for k in range (36): #model num for j in range (31): #latitude for i in range (180): # longitude if t_test_1v1[169,k,j,i]==0: ET[k,j,i]=0 continue elif np.any(t_test_1v1[:,k,j,i]==1): ET_value=max([count for count, item in enumerate(t_test_1v1[1:169,k,j,i]) if item == element], default=0) ET[k,j,i]=ET_value+1921 continue else: ET[k,j,i]=1920Input (xarray DataArray):
<xarray.DataArray (time: 240, deptht: 36, latitude: 31, longitude: 180)> array([[[[0, 0, 1, ..., 1, 1, 1], [0, 1, 1, ..., 0, 0, 0], [1, 1, 0, ..., 0, 0, 1], ..., [0, 0, 0, ..., 0, 0, 0], [0, 0, 0, ..., 0, 0, 0], [0, 0, 0, ..., 0, 0, 0]], [[1, 1, 1, ..., 1, 1, 1], [1, 1, 1, ..., 1, 1, 1], [1, 1, 1, ..., 1, 1, 1], Coordinates: * Time (end_year) datetime64[ns] 1921-12-31 1922-12-31 ... 2100-12-31 * deptht (deptht) int64 1 2 3 4 5 6 7 8 9 ... 28 29 30 31 32 33 34 35 36 * longitude (longitude) float64 30.0 32.0 34.0 36.0 ... 384.0 386.0 388.0 * latitude (latitude) float64 -36.0 -34.0 -32.0 -30.0 ... 32.0 34.0 36.0Expected output (xarray DataArray):
<xarray.DataArray (deptht:36, latitude: 37, longitude: 180)> array([[1983., 2011., 2022., ..., 1937., 1937., 1962.], [2048., 2081., 2083., ..., 1920., 0., 2011.], [2044., 1920., 1993., ..., 0., 0., 1920.], ..., [2004., 1993., 1993., ..., 0., 2010., 2011.], [1920., 1998., 1988., ..., 2011., 2014., 2014.], [2000., 0., 0., ..., 2014., 2011., 2000.]]) Coordinates: * deptht (deptht) int64 1 2 3 4 5 6 7 8 9 ... 28 29 30 31 32 33 34 35 36 * longitude (longitude) float64 30.0 32.0 34.0 36.0 ... 384.0 386.0 388.0 * latitude (latitude) float64 -36.0 -34.0 -32.0 -30.0 ... 32.0 34.0 36.0
Solution 1: Vectorized NumPy Approach
Nested Python loops are slow for large arrays because they don't leverage NumPy's optimized C-backed operations. Instead, we can use fully vectorized functions to compute the result in a single pass:
import numpy as np # Assume t_test_1v1 is your 4D numpy array (time, model, lat, lon) time_length = t_test_1v1.shape[0] # Create an array of years matching each time step: 1921, 1922, ..., 2100 years = np.arange(1921, 1921 + time_length) # Reverse the time axis to turn "last occurrence" into "first occurrence" reversed_data = t_test_1v1[::-1, ...] # Find the first index of 0 in each reversed time series last_zero_reversed_idx = np.argmax(reversed_data == 0, axis=0) # Convert reversed index back to original time axis index original_last_zero_idx = time_length - 1 - last_zero_reversed_idx # Base result: year of the last 0 (defaults to 1921 for all-1 series) result = years[original_last_zero_idx] # Identify special cases all_zero = (t_test_1v1.sum(axis=0) == 0) # Sum is 0 only if all values are 0 all_one = (t_test_1v1.sum(axis=0) == time_length) # Sum equals time length only if all are 1 # Apply the rules to overwrite special cases result[all_zero] = 0 result[all_one] = 1920 # Result is now your desired 3D array (model, lat, lon)
How this works:
- Reverse trick:
np.argmaxreturns the first index where a condition is True. By reversing the time axis, the first 0 we find is the last 0 in the original sequence. - Vectorized checks: Summing along the time axis lets us quickly flag all-0 and all-1 series without looping.
- Minimal overhead: All operations run in optimized C code, so this will be orders of magnitude faster than nested Python loops for large arrays.
Solution 2: Xarray Approach (Preserves Coordinates)
Since your input and output use xarray DataArrays, using xarray's built-in functions will automatically preserve your coordinate metadata, making the code cleaner and less error-prone:
import xarray as xr # Assume t_test_1v1 is your xarray DataArray time_dim = 'time' # Extract integer years from the time coordinate years = t_test_1v1[time_dim].dt.year.values # Reverse the time dimension reversed_da = t_test_1v1.isel({time_dim: slice(None, None, -1)}) # Find the first index of 0 in each reversed time series last_zero_reversed_idx = (reversed_da == 0).argmax(dim=time_dim) # Convert reversed index to original time axis index original_last_zero_idx = len(reversed_da[time_dim]) - 1 - last_zero_reversed_idx # Build the result DataArray with preserved coordinates result_da = xr.DataArray( years[original_last_zero_idx], dims=t_test_1v1.dims[1:], coords={dim: t_test_1v1.coords[dim] for dim in t_test_1v1.dims[1:]} ) # Handle special cases with xarray's where method all_zero = (t_test_1v1.sum(dim=time_dim) == 0) all_one = (t_test_1v1.sum(dim=time_dim) == len(t_test_1v1[time_dim])) result_da = result_da.where(~all_zero, 0) result_da = result_da.where(~all_one, 1920) # result_da is your final 3D xarray DataArray with all original coordinates intact
Why this is better:
- Coordinate safety: You don't have to manually reconstruct coordinate arrays, which eliminates the risk of misalignment.
- Same performance as NumPy: Xarray uses NumPy under the hood, so you get the same speed benefits while working with labeled data.
Performance Note
Both approaches run in O(N) time where N is the total number of elements in the 4D array, but with drastically lower constant factors than Python loops. For large datasets, this should reduce your runtime from minutes/hours to seconds.
内容的提问来源于stack exchange,提问作者Gopika Suresh

