如何利用Python的pool()优化现有for循环以提升处理速度?
多进程处理zonal_stats时Pandas DataFrame为空的问题解决
问题背景
处理5000+普查区的多年每日PRISM气温栅格数据时,单进程循环调用zonal_stats计算均值的总耗时长达两天。改用multiprocessing.Pool改写后,目标填充的Pandas DataFrame data_map始终为空。
原单进程代码:
for n in range(1, len(tract_id_list)): i = 1 for rast in os.listdir(r'/mnt/local_drive/britton/PRISM_data/PRISM_daily_tmax'): if rast[-4: ] == '.tif': tmax = rasterio.open(r'/mnt/local_drive/britton/PRISM_data/PRISM_daily_tmax' + '//' + rast) tmax_array = tmax.read(1) affine = tmax.transform tract_average = zonal_stats(tract_polygon[(n-1):n], tmax_array, affine = affine, stats = ['mean'], all_touched = True, geojason_out = False) tract_average = tract_average[0]['mean'] if n == 1: data_map.loc[i]['Date'] = rast[11:-4] data_map.iloc[(i-1):i, n] = tract_average * 1.8 + 32 # optional conversion to F i = i + 1
改写后的多进程代码(运行后data_map为空):
def spatialaverage(value): i = 1 for rast in os.listdir(r'/mnt/local_drive/britton/PRISM_data/PRISM_daily_tmax'): if rast[-4: ] == '.tif': tmax = rasterio.open(r'/mnt/local_drive/britton/PRISM_data/PRISM_daily_tmax' + '//' + rast) tmax_array = tmax.read(1) affine = tmax.transform tract_average = zonal_stats(tract_polygon[(value-1):value], tmax_array, affine = affine, stats = ['mean'], all_touched = True, geojason_out = False) tract_average = tract_average[0]['mean'] if value == 1: data_map.loc[i]['Date'] = rast[11:-4] data_map.iloc[(i-1):i, value] = tract_average * 1.8 + 32 # optional conversion to F i = i + 1 pool_obj = multiprocessing.Pool() process = pool_obj.map(spatialaverage, range(1, len(tract_id_list))) data_map
问题原因
- 多进程内存隔离:每个子进程会复制主进程的
data_map对象,子进程内对data_map的修改仅作用于本地副本,主进程的原始data_map不会被更新。 - 冗余IO操作:每个子进程都重复遍历、打开并读取所有栅格文件,造成大量重复IO,不仅没提升效率,反而可能拖慢整体速度。
- 索引赋值风险:原代码中
data_map.loc[i]['Date']这种链式索引可能触发SettingWithCopyWarning,导致赋值操作实际未生效。
解决方案
重构代码:子进程返回结果,主进程合并
核心思路是让子进程仅负责计算对应普查区的所有日期均值,返回结果列表;主进程提前预处理所有栅格数据(避免重复IO),再将子进程的结果统一填充到data_map中。
完整代码:
import multiprocessing import os import rasterio from rasterstats import zonal_stats import pandas as pd # 提前预处理所有栅格,避免子进程重复读取 def preprocess_rasters(raster_dir): raster_info = [] for rast in os.listdir(raster_dir): if rast.endswith('.tif'): rast_path = os.path.join(raster_dir, rast) with rasterio.open(rast_path) as src: raster_array = src.read(1) affine = src.transform # 提取日期字符串 date_str = rast[11:-4] raster_info.append((date_str, raster_array, affine)) # 按日期排序,保证所有进程处理顺序一致 raster_info.sort(key=lambda x: x[0]) return raster_info # 单个普查区的计算函数 def compute_tract_stats(tract_idx, raster_info, tract_polygon): # 获取当前普查区的多边形 tract_poly = tract_polygon.iloc[tract_idx:tract_idx+1] stats_list = [] for _, raster_array, affine in raster_info: # 计算区域均值 avg_val = zonal_stats( tract_poly, raster_array, affine=affine, stats=['mean'], all_touched=True, geojson_out=False )[0]['mean'] # 转换为华氏度 stats_list.append(avg_val * 1.8 + 32) return stats_list if __name__ == '__main__': # 配置路径和参数 raster_dir = r'/mnt/local_drive/britton/PRISM_data/PRISM_daily_tmax' tract_count = len(tract_id_list) # 1. 预处理栅格数据 raster_info = preprocess_rasters(raster_dir) dates = [item[0] for item in raster_info] # 2. 初始化data_map data_map = pd.DataFrame( index=range(len(dates)), columns=['Date'] + tract_id_list ) data_map['Date'] = dates # 3. 多进程计算 with multiprocessing.Pool() as pool: # 构造任务参数:每个任务对应一个普查区的索引 tasks = [(idx, raster_info, tract_polygon) for idx in range(tract_count)] # 使用starmap传递多参数 results = pool.starmap(compute_tract_stats, tasks) # 4. 将结果填充到data_map for tract_idx, stats in enumerate(results): data_map.iloc[:, tract_idx + 1] = stats # 查看结果 print(data_map.head())
关键优化说明
- 预加载栅格:在主进程一次性读取所有栅格的数组和变换参数,子进程直接复用,彻底消除重复IO开销。
- 规避内存隔离:子进程返回计算结果,主进程统一赋值,避免直接在子进程中修改共享DataFrame。
- 正确赋值方式:使用
iloc直接赋值,避免链式索引导致的赋值失效问题。 - 灵活参数传递:用
starmap支持多参数任务,比map更适配复杂场景。
额外性能建议
- 限制进程数:在
Pool中指定processes参数(如multiprocessing.Pool(processes=4)),避免进程过多导致CPU资源竞争。 - 分块处理:若栅格数据过大导致内存不足,可分批次读取栅格,而非一次性全部加载。
- 验证多边形索引:确保
tract_polygon.iloc[tract_idx:tract_idx+1]能正确定位到目标普查区的多边形数据。
内容的提问来源于stack exchange,提问作者Annie Britton
相关产品推荐
相关产品推荐

