Python并行计算参数更新:解决Dask实现中参数无法更新问题
迭代计算Dask并行化方案
问题根因
- 迭代步之间存在强依赖:后一步的输入参数完全依赖前一步的计算结果,你将所有迭代步任务同时提交到集群时,后续步骤拿到的参数数组还未被更新,全为初始0值。
- 分布式环境内存不共享:worker进程中修改的只是本地的参数数组副本,主进程和其他worker进程无法感知到修改,最终主进程的参数数组自然只有初始值。
正确实现思路
迭代步本身的串行依赖无法打破,只对每个迭代步内部相互独立的calc_result计算做并行化,既保证计算逻辑正确性,又能拿到并行加速收益。
可运行代码
import numpy as np import random import time from dask.distributed import Client def calc_params(res: list) -> list: time.sleep(random.random()) return [r * 1.1 for r in res] def calc_result(param: float): time.sleep(random.random()) return param + 1 def main(): tic = time.perf_counter() nsteps = 10 nmodels = 4 init_params = [5, 4.5, 8, 2] params = np.zeros((nsteps, nmodels)) params[0] = init_params results = np.zeros((nsteps, nmodels)) for step in range(nsteps): # 并行执行当前步的所有calc_result任务 futures = client.map(calc_result, params[step]) step_results = client.gather(futures) results[step] = step_results if step < nsteps - 1: params[step + 1] = calc_params(step_results) toc = time.perf_counter() print(f'\nElapsed time {toc - tic:.2f} s\n') print(f'Parameters\n{params}\n') print(f'Results\n{results}') if __name__ == '__main__': np.set_printoptions(precision=2) client = Client(n_workers=8) print('\n' + client.dashboard_link) main() client.close()
这个实现输出和串行代码完全一致,运行耗时约为串行版本的1/4,符合并行加速预期。
内容的提问来源于stack exchange,提问作者wigging
相关产品推荐
相关产品推荐

