Python中如何对lambda函数求导生成n阶导对应的lambda函数
实现思路
不要采用逐点数值差分近似的方案,这类方案计算n阶导会引入O(h^n)级别的截断误差,递归计算效率随阶数升高指数下降。直接走符号代数求导路线,生成精确导函数表达式后再封装为lambda,全程无精度损失,高阶导计算效率稳定。
之前用sympy未得到理想效果,通常是未正确提取lambda的表达式结构,直接传入函数对象导致sympy无法解析内部运算逻辑。
可直接复用的实现代码
import inspect import sympy as sp from sympy.utilities.lambdify import lambdify def derivative_lambda(f, n=1, var_name='x'): """ 输入单变量lambda函数f,返回其n阶导对应的lambda函数 参数: f: 输入的纯Python实现单变量lambda,无C扩展调用 n: 求导阶数,默认值为1 var_name: 函数自变量名,默认值为x """ # 提取lambda源码,切分出冒号后的表达式部分 f_source = inspect.getsource(f).strip() expr_str = f_source.split(':', 1)[1].strip() # 定义符号变量 x = sp.symbols(var_name) # 字符串表达式转sympy符号对象,可按需补充local_dict映射math模块函数 expr = sp.parse_expr( expr_str, local_dict={'sin':sp.sin, 'cos':sp.cos, 'tan':sp.tan, 'exp':sp.exp, 'log':sp.log, 'sqrt':sp.sqrt} ) # 计算n阶符号导数 deriv_expr = sp.diff(expr, x, n) # 符号表达式转可直接调用的lambda return lambdify(x, deriv_expr, 'math')
效果验证
# 测试基础案例 f = lambda x : x**2 f1 = derivative_lambda(f, n=1) print(f1(5)) # 输出10,和lambda x:2*x的计算结果完全一致 f2 = derivative_lambda(f, n=2) print(f2(999)) # 输出2,等价于常数函数lambda x:2 # 测试复杂函数 f_test = lambda x : x**4 + 3*x + exp(x) f3_test = derivative_lambda(f_test, n=3) print(f3_test(0)) # 三阶导为24*x + exp(x),代入x=0输出1.0,结果精确
适用边界说明
- 输入lambda必须是纯Python表达式实现的单变量数学函数,无C扩展调用、无分支/循环/IO等非表达式逻辑,完全匹配给定的使用前提
- 求任意阶导数都只需要一次符号解析+一次求导+一次lambda封装,不会因为阶数升高出现效率陡降、误差累积的问题
- 如果lambda的自变量名不是默认的
x,比如自变量是t,调用函数时传入var_name='t'即可适配 - 若表达式中用到更多数学函数,只需要在
parse_expr的local_dict参数里补充对应sympy函数映射即可正常解析
内容的提问来源于stack exchange,提问作者Cardstdani
相关产品推荐
相关产品推荐

