Python使用lambda实现递推逻辑时触发递归深度超限错误
递推函数实现递归超限问题排查
问题背景
给定自然数n,theta、eta为长度为n的正向量,epsilon为长度为n、元素取值为-1或1的向量。需实现算法计算实函数有限序列g=(g_1,...g_n),其中g_n=0,满足如下递推关系:
- 当
x*epsilon_i > x_star * epsilon_i时,g_(i-1)(x)=f_i(x),否则g_(i-1)(x)=0 - 其中
f_i(x)=2eta_i*(x-theta_i)+g_i(x),x_star是f_i的零点(f_i为单调递增的连续分段仿射函数,存在唯一零点)
原实现代码
其中computing_zero为辅助函数,用于在已知f断点的前提下计算f的零点x_star,原代码如下:
def computing_g(theta,epsilon,eta): n=len(theta) g=[lambda x:0,lambda x:2*eta[n-1]*max(0,x-theta[n-1])] # initialization of g : g=[g_n,g_(n-1)] breakpoints_of_f=[theta[n-1]] for i in range(1,n): f= lambda x:2*eta[n-1-i]*(x-theta[n-1-i])+g[i](x) x_star=computing_zero(breakpoints_of_f,f) breakpoints_of_f.append(x_star) g.append(lambda x: f(x) if epsilon[n-1-i]*x > epsilon[n-1-i]*x_star else 0) return(breakpoints_of_f,g)
报错信息
运行代码时触发递归深度超限错误:
line 6, in <lambda> f= lambda x:2*eta[n-1-i]*(x-theta[n-1-i])+g[i](x) line 9, in <lambda> g.append(lambda x: f(x) if epsilon[n-1-i]*x > epsilon[n-1-i]*x_star else 0) RecursionError: maximum recursion depth exceeded in comparison
问题根因
错误不是递推逻辑本身存在无限循环,而是Python闭包的晚绑定机制导致的:
- 循环内用lambda定义函数时,函数体引用的
i、f、x_star、eta[n-1-i]这类和循环迭代相关的变量,不会在lambda定义时就固定为当前轮次的值,而是会在函数实际被调用时,才去查找这些变量的最新取值。 - 当后续调用g序列中的函数时,for循环早已执行完毕,所有循环变量都停留在最后一轮迭代的取值:最后一轮定义的
f会引用g[i],而刚append到g列表的新lambda又会引用这个f,直接形成f调用g[i] -> g[i]调用f的死循环,最终触发递归深度超限。
修复方案
解决思路是在定义lambda时,把每一轮用到的循环变量通过函数默认参数的形式固定——Python函数的默认参数在函数定义阶段就完成求值,不会随外部变量变动,可以完美避开晚绑定的坑。
修复后代码如下:
def computing_g(theta,epsilon,eta): n = len(theta) # 初始化g序列,初始化项同样固定参数避免外部变量影响 g = [ lambda x: 0, lambda x, eta_val=eta[n-1], th=theta[n-1]: 2 * eta_val * max(0, x - th) ] breakpoints_of_f = [theta[n-1]] for i in range(1, n): curr_idx = n - 1 - i # 提前取出当前轮次的固定参数 curr_eta = eta[curr_idx] curr_theta = theta[curr_idx] curr_eps = epsilon[curr_idx] prev_g = g[i] # 定义f时固定当前轮参数 f = lambda x, eta_val=curr_eta, th=curr_theta, g_i=prev_g: 2 * eta_val * (x - th) + g_i(x) x_star = computing_zero(breakpoints_of_f, f) breakpoints_of_f.append(x_star) # 定义新的g项时固定当前轮的f、epsilon、x_star g.append( lambda x, f_ref=f, eps=curr_eps, xs=x_star: f_ref(x) if eps * x > eps * xs else 0 ) return (breakpoints_of_f, g)
内容的提问来源于stack exchange,提问作者Skywear
相关产品推荐
相关产品推荐

