Python中含NaN的[time,lat,lon]三维数组Pearson相关计算优化
Hey there! Let's fix that slow grid-wise correlation code of yours—nested loops over lat/lon are brutal for performance, and that empty result from stack() is almost certainly due to how NaNs are being handled. Here's how to get this right and fast:
1. Full Time-Series Correlation (No Threshold)
Xarray's built-in xr.corr function is made exactly for this use case! It automatically computes correlations along a specified dimension (here, time), skips NaNs by default, and avoids slow Python-level loops entirely.
# Read and slice your data first a = xr.open_dataset(A).sel(time=slice('1950-03','2013-12'))[var1] b = xr.open_dataset(B).sel(time=slice('1950-03','2013-12'))[var2] # Compute full correlation in one line corr_full = xr.corr(a, b, dim='time')
This runs way faster than your original nested loops and correctly ignores NaNs in each grid's time series.
2. Threshold-Masked Correlation (Only when A > TH)
For this, we need to filter each grid's time series to only include points where a > TH (and exclude any NaNs in either series). The best tool here is xr.apply_ufunc, which lets us apply a custom correlation function to each grid's time series efficiently.
We'll use scipy.stats.pearsonr because it's robust and returns the correlation coefficient directly. We'll also add a check to return NaN if there are fewer than 2 valid points (you can't compute a correlation with 1 or 0 samples!).
import scipy.stats def thresholded_corr(a_ts, b_ts, thresh): # Create a mask for valid points: a > threshold, and no NaNs in either series valid_mask = (a_ts > thresh) & ~np.isnan(a_ts) & ~np.isnan(b_ts) # Skip grids with too few valid samples to compute correlation if valid_mask.sum() < 2: return np.nan # Compute Pearson correlation on the filtered time series return scipy.stats.pearsonr(a_ts[valid_mask], b_ts[valid_mask])[0] # Apply the function to every grid cell efficiently corr_th = xr.apply_ufunc( thresholded_corr, a, b, input_core_dims=[['time'], ['time']], # Tell Xarray which dims are per-grid time series output_core_dims=[[]], # Output is a scalar value per grid kwargs={'thresh': TH}, vectorize=True, # Automatically loop over lat/lon (optimized C-level loops) dask='allowed' # Use this if your data is chunked with Dask for parallel processing )
Putting It All Together (Your Updated Function)
Here's how to replace your original slow function with these optimized methods:
import xarray as xr import numpy as np import scipy.stats def correlate(A, B, var1, var2, TH): output_filename = f"corr_{var1}_{var2}_TH_{TH}.nc" # Load and slice your datasets to the desired time range a = xr.open_dataset(A).sel(time=slice('1950-03','2013-12'))[var1] b = xr.open_dataset(B).sel(time=slice('1950-03','2013-12'))[var2] # 1. Compute full time-series correlation corr_full = xr.corr(a, b, dim='time') # 2. Compute threshold-masked correlation def thresholded_corr(a_ts, b_ts, thresh): valid_mask = (a_ts > thresh) & ~np.isnan(a_ts) & ~np.isnan(b_ts) if valid_mask.sum() < 2: return np.nan return scipy.stats.pearsonr(a_ts[valid_mask], b_ts[valid_mask])[0] corr_th = xr.apply_ufunc( thresholded_corr, a, b, input_core_dims=[['time'], ['time']], output_core_dims=[[]], kwargs={'thresh': TH}, vectorize=True, dask='allowed' ) # Build the output dataset with both correlation fields ds = xr.Dataset( { 'r': corr_full, f'r_TH_{TH}': corr_th }, coords={'lon': a['lon'], 'lat': a['lat']} ) # Optional: Save directly to NetCDF if needed # ds.to_netcdf(output_filename) return ds
Why Your Original stack() Method Failed
When you used stack() and computed mean/std directly, any NaN in a grid's time series made the mean/std return NaN—this turned your entire covariance and correlation arrays into NaNs. To fix that stack-based approach for full correlations, you'd need to use np.nanmean and np.nanstd (which skip NaNs):
# Alternative stack-based method for full correlation (for reference) A_stack = a.stack(z=('lat','lon')) B_stack = b.stack(z=('lat','lon')) # Use nanmean/nanstd to ignore NaNs in each grid's time series A_mean = np.nanmean(A_stack, axis=0) B_mean = np.nanmean(B_stack, axis=0) A_std = np.nanstd(A_stack, axis=0, ddof=1) # ddof=1 for sample standard deviation B_std = np.nanstd(B_stack, axis=0, ddof=1) cov = np.nanmean((A_stack - A_mean) * (B_stack - B_mean), axis=0) corr_stack = cov / (A_std * B_std) # Unstack back to lat/lon dimensions corr_stack = corr_stack.unstack('z')
Note that this only works for full correlations—for thresholded cases, apply_ufunc is still the most flexible and clean approach.
内容的提问来源于stack exchange,提问作者David Hoffmann

