scipy.optimize.minimize中tuple与np.array结果差异问题求助
问题根源与解决办法
问题原因
你的代码核心问题出在参数传递逻辑:
- 传入
tuple时,minimize的args=generate会把tuple的每个元素拆成单独参数,被目标函数的*sales接收,此时sales是所有数据点的序列,函数内循环逻辑正常执行,但遍历大量单独参数会拖慢速度。 - 传入
np.array时,args=generate会把整个数组作为单个参数传入,*sales此时只包含一个元素(整个数组),导致函数内range(1, len(sales))的循环根本不执行,返回的误差始终为0,最终得到错误的优化结果。
解决方案
修改目标函数,让它直接接受一个序列(tuple/np.array)作为输入,不再使用可变参数*sales。同时利用numpy的向量化运算进一步提速,避免Python循环的性能损耗。
修改后的代码
import numpy as np from scipy.optimize import minimize import time import warnings warnings.filterwarnings("ignore", category=RuntimeWarning) warnings.filterwarnings("ignore", category=UserWarning) start_time = time.time() def Bass1(x, P, Q, M): return (P * M + (Q - P) * x) - (Q / M) * (x ** 2) def squareMistake1(k, sales) -> float: P, Q, M = k c0 = sales[0] cumulative = np.zeros_like(sales) cumulative[0] = c0 for i in range(1, len(sales)): p = Bass1(cumulative[i-1], P, Q, M) cumulative[i] = cumulative[i-1] + p # 批量计算平方误差和 return np.sum((cumulative - sales) ** 2) # 准备数据与参数 k0 = [0.0008791696672306727, 0.19252826585535315, 3328.9193848309856] kb = ((0, None), (0, None), (0, None)) # 使用numpy数组(速度快,结果正确) generate = np.array([8.26192344363636, 9.20460066059596, 12.0178164697778, 15.921260267805, 21.2161740066094, 31.420434564131, 38.3904519471421, 52.3307819867071, 62.9113953016839, 85.1161924282732, 104.083879757882, 132.859216030029, 170.682620580279, 220.600045153997, 276.020526299077, 346.465021938078, 440.385091980306, 530.55442135112, 635.49205101167, 705.805860788812, 831.42968828187, 962.227395409379, 1140.31094904253, 1269.52053571083, 1418.17004626655, 1591.2135122193]) method_list = ['Nelder-Mead', 'Powell', 'L-BFGS-B', 'TNC', 'SLSQP', 'trust-constr'] for method in method_list: try: res = minimize(squareMistake1, k0, args=(generate,), method=method, bounds=kb) print(f"方法 {method}: k = {tuple(res.x)}") except Exception as e: print(f"方法 {method} 失败: {str(e)}") end_time = time.time() print(f"执行时间: {end_time - start_time:.2f} 秒")
关键修改点
- 目标函数参数调整:将
*sales改为sales,直接接收序列输入,调用minimize时通过args=(generate,)传递数组(末尾逗号保证传入的是元组)。 - 误差计算优化:用numpy数组存储累积值,最后用
np.sum批量计算平方误差和,比Python循环累加更快。 - 逻辑一致性:无论传入
tuple还是np.array,函数内部处理逻辑完全一致,既保证结果正确,又能利用numpy的高效运算提升速度。
效果验证
修改后使用np.array作为输入,执行速度与原代码的np.array版本一致,同时优化结果与原代码tuple版本完全匹配,解决了速度与正确性的矛盾。
内容的提问来源于stack exchange,提问作者Михаил Никифоров
相关产品推荐
相关产品推荐

