如何在scipy differential evolution多worker下传递共享变量跟踪实验编号
问题:scipy differential_evolution多worker模式下实验编号无法连续递增
在使用scipy的differential_evolution函数时,设置workers=1可正常跟踪连续的实验编号,但workers>1时,实验编号会出现重复、不连续的问题。
原始代码
from scipy.optimize import differential_evolution import sys def objectiveFunction(x, experiment_no): y = x[0] + x[1] # 跟踪实验编号 experiment_no["exp"] = experiment_no["exp"] + 1 print(f"This is experiment no : {experiment_no['exp']}") return y if __name__ == "__main__": minRanges = [1.2, 10] maxRanges = [1.5, 20] experiment_no = {"exp": 0} try: bounds = list(zip(minRanges, maxRanges)) result = differential_evolution( objectiveFunction, bounds, args=(experiment_no,), strategy="best1bin", workers=1, maxiter=2, updating="deferred", polish=False ) print('Global minimum [x]:') print(result.x) print('Function value at global minimum [f(x)]:') print(result.fun) except: exit(sys.exc_info()[:2])
现象对比
workers=1时,输出为连续递增的编号:
This is experiment no :1 This is experiment no :2 This is experiment no :3 ...
workers=2时,输出出现重复、混乱:
This is experiment no :1 This is experiment no :1 This is experiment no :11 ...
原因分析
当workers>1时,differential_evolution默认使用多进程执行目标函数。普通Python字典属于进程私有数据,每个子进程会复制一份初始的experiment_no字典,各进程修改的都是自己的副本,不会同步到主进程或其他子进程,因此出现编号重复、不连续的问题。
解决方案
使用multiprocessing模块提供的共享变量(如Value)跟踪实验编号,同时通过锁机制保证递增操作的原子性,避免多进程同时修改导致的竞态条件。
修改后的代码
from scipy.optimize import differential_evolution import multiprocessing import sys def objectiveFunction(x, experiment_no): y = x[0] + x[1] # 加锁保证原子操作,避免多进程同时修改 with experiment_no.get_lock(): experiment_no.value += 1 current_exp = experiment_no.value print(f"This is experiment no : {current_exp}") return y if __name__ == "__main__": minRanges = [1.2, 10] maxRanges = [1.5, 20] # 创建进程共享的整数变量,初始值为0('i'表示整数类型) experiment_no = multiprocessing.Value('i', 0) try: bounds = list(zip(minRanges, maxRanges)) result = differential_evolution( objectiveFunction, bounds, args=(experiment_no,), strategy="best1bin", workers=2, maxiter=2, updating="deferred", polish=False ) print('Global minimum [x]:') print(result.x) print('Function value at global minimum [f(x)]:') print(result.fun) except: exit(sys.exc_info()[:2])
关键说明
- 共享变量创建:
multiprocessing.Value('i', 0)创建进程间共享的整数变量,'i'指定变量类型为整数,初始值设为0。 - 锁机制:
with experiment_no.get_lock()上下文管理器确保experiment_no.value += 1是原子操作,避免多进程同时修改导致计数错误。 - 变量访问:共享变量通过
.value属性访问和修改,替代原有的字典键值方式。
测试结果
设置workers=2时,输出将变为连续递增的实验编号:
This is experiment no :1 This is experiment no :2 This is experiment no :3 This is experiment no :4 ...
内容的提问来源于stack exchange,提问作者First Python
相关产品推荐
相关产品推荐

