Python优化:非结构化点云转GeoTIFF的凸域内数组插值提速
Python点云插值优化:大尺寸含NaN数组生成GeoTIFF
问题背景
从非结构化点云生成GeoTIFF时,现有10000×9300的2D数组包含大量NaN值(需填充的空白区域)。使用scipy.interpolate.griddata单文件处理耗时约15分钟,上百个文件的总耗时完全无法接受。约束条件:禁止对项目区域凸域外的NaN点进行外推插值。
现有实现代码
import numpy as np from scipy.interpolate import griddata zi = np.load("Array.npy") x, y = np.indices(zi.shape) zi_i = griddata( (x[~np.isnan(zi)], y[~np.isnan(zi)]), zi[~np.isnan(zi)], (x[np.isnan(zi)], y[np.isnan(zi)]) )
优化方案
1. 过滤凸域外NaN点,减少计算量
先计算有效数据点的凸包,仅对凸包内部的NaN点执行插值,直接削减无意义的计算量:
import numpy as np from scipy.interpolate import griddata from scipy.spatial import ConvexHull zi = np.load("Array.npy") x, y = np.indices(zi.shape) # 提取有效点坐标与值 valid_mask = ~np.isnan(zi) valid_points = np.column_stack((x[valid_mask], y[valid_mask])) valid_vals = zi[valid_mask] # 计算有效点凸包 hull = ConvexHull(valid_points) # 批量判断NaN点是否在凸包内 nan_coords = np.column_stack((x[np.isnan(zi)], y[np.isnan(zi)])) in_hull_mask = np.all( np.dot(hull.equations[:, :-1], nan_coords.T) + hull.equations[:, -1] <= 1e-12, axis=0 ) # 仅对凸包内的NaN点插值 zi_i = zi.copy() target_coords = nan_coords[in_hull_mask] if len(target_coords) > 0: interp_vals = griddata(valid_points, valid_vals, target_coords, method="linear") # 将插值结果回填到原数组 zi_i[np.isnan(zi)][in_hull_mask] = interp_vals
2. 替换为高效插值工具
方案A:GDAL栅格插值(C++底层,速度提升显著)
GDAL的gdal.Grid专为栅格插值优化,支持多种算法,性能远优于scipy:
from osgeo import gdal, osr import numpy as np zi = np.load("Array.npy").astype(np.float32) # 创建内存GDAL数据集 driver = gdal.GetDriverByName('MEM') src_ds = driver.Create('', zi.shape[1], zi.shape[0], 1, gdal.GDT_Float32) src_ds.GetRasterBand(1).WriteArray(zi) # 设置地理变换与投影(根据实际数据调整) src_ds.SetGeoTransform([0, 1, 0, zi.shape[0], 0, -1]) srs = osr.SpatialReference() srs.ImportFromEPSG(4326) src_ds.SetProjection(srs.ExportToWkt()) # 执行线性插值,自动忽略凸域外的NaN dst_ds = gdal.Grid( '/vsimem/result.tif', src_ds, algorithm='linear', noData=np.nan ) # 读取插值结果 zi_i = dst_ds.GetRasterBand(1).ReadAsArray()
方案B:Numba加速插值逻辑
用Numba JIT编译核心插值代码,针对数组操作提速:
import numpy as np from numba import jit from scipy.spatial import cKDTree zi = np.load("Array.npy") x, y = np.indices(zi.shape) valid_mask = ~np.isnan(zi) valid_points = np.column_stack((x[valid_mask], y[valid_mask])) valid_vals = zi[valid_mask] # 构建KD-Tree快速查找近邻 tree = cKDTree(valid_points) nan_coords = np.column_stack((x[np.isnan(zi)], y[np.isnan(zi)])) # 查询每个NaN点的3个近邻(线性插值需求) distances, indices = tree.query(nan_coords, k=3) # Numba加速线性插值计算 @jit(nopython=True, parallel=True) def numba_linear_interp(valid_vals, indices, distances): interp_vals = np.full(len(indices), np.nan) for i in range(len(indices)): idx = indices[i] d = distances[i] # 避免除以0,跳过距离为0的点 if d[0] < 1e-12: interp_vals[i] = valid_vals[idx[0]] continue # 线性插值权重计算 weights = 1 / d weights /= weights.sum() interp_vals[i] = np.sum(valid_vals[idx] * weights) return interp_vals # 执行插值并回填 interp_vals = numba_linear_interp(valid_vals, indices, distances) zi_i = zi.copy() zi_i[np.isnan(zi)] = interp_vals
3. 多进程批量处理文件
针对上百个文件,用多进程充分利用CPU多核:
import multiprocessing as mp import numpy as np def process_file(file_path): # 这里嵌入单个文件的插值逻辑(结合上述优化方案) zi = np.load(file_path) # ... 插值处理 ... # 保存为GeoTIFF(需自行实现save_geotiff函数) save_geotiff(zi_i, f"{file_path.split('.')[0]}_output.tif") if __name__ == "__main__": file_list = ["file1.npy", "file2.npy", ...] # 你的文件列表 with mp.Pool(mp.cpu_count()) as pool: pool.map(process_file, file_list)
4. 内存优化
- 精度允许时,将数组从
float64转为float32:zi = np.load("Array.npy").astype(np.float32),内存占用直接减半。 - 分块处理:将大数组拆分为多个小块,逐块插值后拼接,避免内存瓶颈。
内容的提问来源于stack exchange,提问作者Victor Alteirac
相关产品推荐
相关产品推荐

