多进程训练预测模型初期快后期骤慢的问题排查求助
多进程并行训练预测模型时性能骤降问题排查
问题背景
我需要在多个回测日期和模型参数上训练预测模型,基于约10年季度数据(40个数据点)对ARIMA、ETS等单变量/多变量模型结果取平均,目标是并行运行数千种参数组合。
自定义模型函数
def train_test_func(model_params): data = read_data_from_pickle() data_train, data_test = train_test_split(data, backtestdate) model1 = ARIMA.fit(data_train) data_pred1 = model1.predict(len(data_test)) # 其他模型训练与预测逻辑 ... results = error_eval(data_pred1, ..., data_pred_i, data_test) save_to_aws_s3(results) logger.info("log steps here")
多进程执行脚本
# 导入自定义训练测试函数 from my_custom_model import train_test_func commands = [] if __name__ == '__main__': # 生成所有参数组合 for backtest_date in target_backtest_dates: for param_a in target_drugs: for param_b in param_b_options: for param_c in param_c_options: args = { "backtest_date": backtest_date, "param_a": param_a, "param_b": param_b, "param_c": param_c } commands.append(args) # 初始化进程池并执行 count = multiprocessing.cpu_count() with multiprocessing.get_context("spawn").Pool(processes=count) as pool: pool.map(train_test_func, batched_args)
性能异常表现
前200次迭代速度较快(约每分钟50次),之后骤降至约每分钟1次;单核心运行反而能达到约每分钟5次。所有进程相互独立,仅使用小型数据集且无依赖。
性能截图说明:
- 第一张截图显示前期处理速度稳定,后期出现显著性能下滑
- 第二张截图展示CPU使用率后期出现异常波动,利用率无法维持高位
可能的问题与解决方向
1. 进程资源累积与泄漏
spawn模式下每个子进程都会重新初始化Python环境,若train_test_func中存在未释放的资源(如未关闭的文件句柄、内存中残留的大对象),随着进程重复执行,会导致系统内存、文件描述符耗尽,引发系统分页(swap)或资源竞争,拖慢整体速度。
解决方法:
- 在
train_test_func末尾显式清理资源,比如删除大变量、关闭文件连接 - 限制进程池大小,不要直接使用
cpu_count(),建议设置为cpu_count() - 1,预留资源给系统进程 - 启用
maxtasksperchild参数,让每个进程执行一定任务后自动重启,避免资源累积:with multiprocessing.get_context("spawn").Pool(processes=count, maxtasksperchild=100) as pool: pool.map(train_test_func, commands)
2. AWS S3写入瓶颈
大量进程同时向S3写入结果,会触发S3的请求限流、网络拥堵,前期任务少竞争小速度快,后期任务集中提交时,每个进程都要等待上传完成,导致整体速度骤降。
解决方法:
- 改为批量上传:每个进程先将结果暂存到本地文件,积累到一定数量后再批量上传到S3
- 单独用线程处理S3上传,让训练逻辑和IO操作并行,避免训练进程被IO阻塞
- 确保计算资源与S3存储在同一区域,减少网络传输延迟
3. 重复读取数据的开销
每个子进程都调用read_data_from_pickle()读取数据,若pickle文件较大,会导致大量重复的磁盘IO和内存占用,后期内存不足引发系统swap,严重拖慢速度。
解决方法:
- 主进程提前读取数据,通过
initializer传递给所有子进程,避免重复读取:def init_worker(shared_data): global data data = shared_data if __name__ == '__main__': # 主进程一次性读取数据 global_data = read_data_from_pickle() with multiprocessing.get_context("spawn").Pool( processes=count, initializer=init_worker, initargs=(global_data,) ) as pool: pool.map(train_test_func, commands) - 子进程的
train_test_func直接使用全局的data变量,无需重复读取
4. 参数传递错误
脚本中生成的参数列表是commands,但最后pool.map传入的是batched_args,若batched_args未正确定义或与commands不一致,会导致进程执行异常,出现隐性等待。
解决方法:
- 确认参数列表正确性,将
batched_args替换为commands,确保每个进程收到正确的参数字典
5. 日志系统的锁竞争
大量进程同时写入同一日志文件,会触发日志库的同步锁竞争,每个进程都要等待锁释放才能写入日志,导致等待时间越来越长。
解决方法:
- 改为每个进程写入独立的本地日志文件,后期再合并
- 使用异步日志框架(如
logging.handlers.QueueHandler),避免同步阻塞
内容的提问来源于stack exchange,提问作者Maximus1990
相关产品推荐
相关产品推荐

