如何向量化网格空间特征提取?卫星数据批量网格取值优化
解决方案:向量化批量提取网格子区域
核心思路
你的问题本质是批量从有序网格中提取矩形子区域,当前的循环推导式在处理7万+特征时效率极低。我们可以利用网格的有序性(lat/lon为线性生成且已排序),结合numpy的向量化索引替代循环,主要分两步优化:
步骤1:用np.searchsorted替代min_diff(更快的最近点查找)
原min_diff用argmin遍历所有点找最近值,效率低下。由于lat和lon是严格有序数组,用np.searchsorted可直接定位插入位置,再比较相邻点找到最近值,速度提升数个数量级:
def find_nearest(sorted_arr, target_vals): # sorted_arr是有序数组,target_vals是待匹配的批量值 pos = np.searchsorted(sorted_arr, target_vals, side='left') # 处理边界:避免索引越界 pos = np.clip(pos, 1, len(sorted_arr)-1) # 比较相邻点距离,选择更近的那个 left_dist = abs(sorted_arr[pos-1] - target_vals) right_dist = abs(sorted_arr[pos] - target_vals) return sorted_arr[np.where(left_dist <= right_dist, pos-1, pos)]
步骤2:向量化批量提取子区域
不再循环每个特征,先把网格数据转换成二维数组,再用所有特征的边界索引范围批量提取:
修改后的main函数
import numpy as np import pandas as pd idx = pd.IndexSlice random = np.random.randint def make_gs(x1, y1, x2, y2, x_size, y_size): col = pd.Index(np.linspace(x1, x2, x_size, dtype=np.float32), name="lon") ndx = pd.Index(np.linspace(y1, y2, y_size, dtype=np.float32), name="lat") grd = ( pd.DataFrame(columns=col, index=ndx).unstack("lat").reset_index().dropna(axis=1) ) grd["wv"] = random(0, 255, size=len(grd)) grd["ir"] = random(0, 255, size=len(grd)) return grd.set_index(["lat", "lon"]).sort_index(level=["lat", "lon"]) def find_nearest(sorted_arr, target_vals): pos = np.searchsorted(sorted_arr, target_vals, side='left') pos = np.clip(pos, 1, len(sorted_arr)-1) left_dist = abs(sorted_arr[pos-1] - target_vals) right_dist = abs(sorted_arr[pos] - target_vals) return sorted_arr[np.where(left_dist <= right_dist, pos-1, pos)] def main(): gs = make_gs(-129, 54, -60, 20, x_size=972, y_size=635) # 提取有序的lat/lon数组 lat_vals = gs.index.unique('lat').to_numpy() lon_vals = gs.index.unique('lon').to_numpy() n_features = 76_020 f = pd.DataFrame( [ {"minx": -126, "maxx": -68, "miny": 37, "maxy": 39}, {"minx": -91, "maxx": -70, "miny": 31, "maxy": 37}, {"minx": -124, "maxx": -64, "miny": 24, "maxy": 26}, ] * (n_features // 3) ) # 批量查找所有边界对应的最近网格点 minx = find_nearest(lon_vals, f['minx'].to_numpy()) maxx = find_nearest(lon_vals, f['maxx'].to_numpy()) miny = find_nearest(lat_vals, f['miny'].to_numpy()) maxy = find_nearest(lat_vals, f['maxy'].to_numpy()) # 把网格数据转换成lat为行、lon为列的二维结构,避免重复索引开销 gs_unstacked = gs.unstack('lon') # 把边界值转换成数组索引(切片左闭右开,所以max索引+1) def get_slice_indices(sorted_arr, target_vals): return np.searchsorted(sorted_arr, target_vals, side='left') y1_idx = get_slice_indices(lat_vals, miny) y2_idx = get_slice_indices(lat_vals, maxy) + 1 x1_idx = get_slice_indices(lon_vals, minx) x2_idx = get_slice_indices(lon_vals, maxx) + 1 # 批量提取子区域:用iloc切片替代loc标签查找,速度提升显著 result = [] for y1, y2, x1, x2 in zip(y1_idx[:100], y2_idx[:100], x1_idx[:100], x2_idx[:100]): # 切片后转回原MultiIndex格式 subset = gs_unstacked.iloc[y1:y2, x1:x2].stack('lon').sort_index() result.append(subset) # 如果不需要DataFrame格式,直接用numpy数组提取(速度极致): # wv_data = gs_unstacked['wv'].values # ir_data = gs_unstacked['ir'].values # wv_subsets = [wv_data[y1:y2, x1:x2] for y1,y2,x1,x2 in zip(y1_idx[:100], y2_idx[:100], x1_idx[:100], x2_idx[:100])] return result if __name__ == "__main__": main()
关键优化点
- 有序数组高效查找:用
np.searchsorted替代argmin,时间复杂度从O(n*m)降至O(m log n)(n为网格点数,m为特征数)。 - 数组切片替代标签查找:
iloc基于位置的切片比loc基于标签的查找快数倍,尤其适合批量处理。 - 避免重复计算:提前将网格数据转换为二维结构,减少pandas索引的重复开销。
内容的提问来源于stack exchange,提问作者Jason Leaver
相关产品推荐
相关产品推荐

