You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何利用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 11:05:18