Python黄金分割搜索算法仅保留打印语句时才返回正确极小值的异常问题排查
问题背景
我正在迁移一段黄金分割搜索算法的Python代码,用于寻找函数极小值,但遇到了一个诡异的问题:只有取消注释两个打印语句时,算法才能返回正确的极小值;而完全相同的算法翻译成Fortran后运行正常。
异常表现:
- 无打印语句时,返回值为初始中点值(疑似循环未充分执行),对应目标函数值并非极小值
- 保留打印语句时,返回结果与Fortran版本一致,为正确极小值
可能原因分析
1. 预先计算迭代次数n的浮点数精度误差
算法中通过n = int(np.ceil(np.log(tol/dist)/np.log(inv_phi)))预先计算迭代次数,但浮点数的舍入误差可能导致n偏小,循环次数不足,无法收敛到正确结果。而打印语句触发的额外函数调用,可能因目标函数的微小计算差异,改变了初始的yc < yd判断,进而让循环走向正确的收敛路径。
2. 目标函数的隐性副作用
如果目标函数或其依赖函数存在未被察觉的副作用(比如全局参数par的属性被意外修改,或par的属性是动态延迟加载的),两次调用obj(c,*args)可能返回不同结果,导致有无打印语句时的初始判断逻辑完全不同。
3. 浮点数比较的精度敏感问题
当yc和yd的值非常接近时,浮点数计算的微小误差可能导致yc < yd的判断结果反转,使得循环分支选择完全不同,最终得到差异极大的优化结果。
解决方法
方法1:替换预先迭代次数为while循环(推荐)
避免依赖n的计算,直接根据区间长度dist是否小于容忍度tol来控制循环,确保迭代足够多次直到收敛:
def optimizer(obj, a, b, args=(), tol=1e-6): """ golden section search optimizer Args: obj (callable): 1d function to optimize over a (double): minimum of starting bracket b (double): maximum of starting bracket args (tuple): additional arguments to the objective function tol (double,optional): tolerance Returns: (float): optimization result """ inv_phi = (np.sqrt(5) - 1) / 2 # 1/phi inv_phi_sq = (3 - np.sqrt(5)) / 2 # 1/phi^2 # a. distance dist = b - a if dist <= tol: return (a+b)/2 # b. potential new mid-points c = a + inv_phi_sq * dist d = a + inv_phi * dist yc = obj(c,*args) yd = obj(d,*args) # d. loop until dist <= tol while dist > tol: if yc < yd: b = d d = c yd = yc dist = inv_phi*dist c = a + inv_phi_sq * dist yc = obj(c,*args) else: a = c c = d yc = yd dist = inv_phi*dist d = a + inv_phi * dist yd = obj(d,*args) # e. return if yc < yd: return (a+d)/2 else: return (c+b)/2
方法2:验证目标函数的纯函数特性
确保目标函数是纯函数(相同输入返回相同输出),可以通过多次调用同一输入并对比结果验证:
import numpy as np # 定义黄金分割常数 inv_phi = (np.sqrt(5) - 1) / 2 inv_phi_sq = (3 - np.sqrt(5)) / 2 # 测试目标函数的一致性 test_c = d_low + inv_phi_sq*(d_high - d_low) test_d = d_low + inv_phi*(d_high - d_low) yc1 = obj_last_period(test_c, x, par) yc2 = obj_last_period(test_c, x, par) yc3 = obj_last_period(test_c, x, par) print(f"yc多次调用结果:{yc1}, {yc2}, {yc3}") print(f"是否一致:{np.isclose(yc1, yc2) and np.isclose(yc1, yc3)}") yd1 = obj_last_period(test_d, x, par) yd2 = obj_last_period(test_d, x, par) yd3 = obj_last_period(test_d, x, par) print(f"yd多次调用结果:{yd1}, {yd2}, {yd3}") print(f"是否一致:{np.isclose(yd1, yd2) and np.isclose(yd1, yd3)}")
如果结果不一致,检查par的属性是否被意外修改,或是否存在动态计算的属性。
方法3:增加浮点数比较的容错
在比较yc和yd时,加入极小的epsilon,避免因微小精度差异导致判断反转:
# 修改循环中的判断条件 if yc < yd - 1e-12: # 加入容错epsilon b = d d = c yd = yc dist = inv_phi*dist c = a + inv_phi_sq * dist yc = obj(c,*args) else: a = c c = d yc = yd dist = inv_phi*dist d = a + inv_phi * dist yd = obj(d,*args)
方法4:使用更高精度的浮点数
尝试将所有计算转换为更高精度的浮点数(比如numpy.float128),减少精度误差:
inv_phi = np.float128((np.sqrt(5) - 1) / 2) inv_phi_sq = np.float128((3 - np.sqrt(5)) / 2) a = np.float128(a) b = np.float128(b)
验证结果
使用方法1的while循环版本,应该能在不添加打印语句的情况下,得到与Fortran版本一致的正确结果,因为它直接根据区间长度判断停止迭代,避免了预先计算n带来的误差。
内容的提问来源于stack exchange,提问作者Bob Millard

