求精准导数计算方案:现有Lua实现存在功能缺陷需优化
优化Lua导数计算的方案与实现
现有实现的问题分析
第一种幂函数专用实现
你的第一个实现完全依赖幂函数求导公式,存在明显局限性:
- 仅支持形如
xⁿ的函数,对eˣ、aˣ、sin(x)等函数完全失效 - 依赖
math.log(Function(x), x)推导指数n,但当Function(x)不是x的幂函数时,这个推导逻辑错误 - 当
x ≤ 0时,math.log会抛出错误,鲁棒性极差
第二种数值微分实现
第二种中心差分的思路是对的,但存在精度优化空间:
- 固定步长无法适配不同大小的
x(比如x=1e6和x=1e-6用同一个步长,误差差异极大) - 手动加入的
round判断可能引入不必要的误差,尤其是在导数为非整数的场景
优化方案与代码实现
方案1:自适应步长的数值微分
通过动态调整步长,平衡截断误差和舍入误差,同时使用标准的中心差分公式提升精度:
function derivative_num(fun, x) -- 基于双精度浮点数的epsilon计算自适应步长 local eps = math.sqrt(2.220446049250313e-16) local h = eps * (1 + math.abs(x)) -- 步长随x的绝对值动态调整 local x_plus = x + h local x_minus = x - h local f_plus = fun(x_plus) local f_minus = fun(x_minus) -- 中心差分:二阶精度,比单侧差分更准确 return (f_plus - f_minus) / (2 * h) end -- 测试用例 print(derivative_num(function(x) return math.exp(x) end, 1)) -- 近似e≈2.71828 print(derivative_num(function(x) return x^3 end, 10)) -- 精确值300 print(derivative_num(function(x) return math.sin(x) end, math.pi/2)) -- 近似0
方案2:符号微分(适用于已知表达式的函数)
通过构建函数的抽象语法树,应用求导规则生成精确的导函数,再代入数值计算:
-- 构建符号表达式 local function sym_exp(op, ...) return {type = op, args = {...}} end -- 符号求导核心逻辑 local function sym_diff(expr, var) if type(expr) == "number" then return 0 elseif expr.type == "var" then return expr.name == var and 1 or 0 elseif expr.type == "pow" then local base, exp = expr.args[1], expr.args[2] -- 处理xⁿ、aˣ、u^v三种幂函数场景 if type(exp) == "number" then return sym_exp("mul", exp, sym_exp("pow", base, exp-1), sym_diff(base, var)) elseif type(base) == "number" then return sym_exp("mul", expr, sym_exp("ln", base), sym_diff(exp, var)) else return sym_exp("mul", expr, sym_exp("add", sym_exp("mul", sym_diff(exp, var), sym_exp("ln", base)), sym_exp("mul", exp, sym_exp("div", sym_diff(base, var), base)) )) end elseif expr.type == "exp" then -- e^x local arg = expr.args[1] return sym_exp("mul", expr, sym_diff(arg, var)) elseif expr.type == "sin" then local arg = expr.args[1] return sym_exp("mul", sym_exp("cos", arg), sym_diff(arg, var)) elseif expr.type == "add" then local a, b = expr.args[1], expr.args[2] return sym_exp("add", sym_diff(a, var), sym_diff(b, var)) elseif expr.type == "mul" then local a, b = expr.args[1], expr.args[2] return sym_exp("add", sym_exp("mul", sym_diff(a, var), b), sym_exp("mul", a, sym_diff(b, var)) ) end end -- 符号表达式求值 local function sym_eval(expr, env) if type(expr) == "number" then return expr end if expr.type == "var" then return env[expr.name] end if expr.type == "pow" then return sym_eval(expr.args[1], env) ^ sym_eval(expr.args[2], env) end if expr.type == "exp" then return math.exp(sym_eval(expr.args[1], env)) end if expr.type == "sin" then return math.sin(sym_eval(expr.args[1], env)) end if expr.type == "cos" then return math.cos(sym_eval(expr.args[1], env)) end if expr.type == "add" then return sym_eval(expr.args[1], env) + sym_eval(expr.args[2], env) end if expr.type == "mul" then return sym_eval(expr.args[1], env) * sym_eval(expr.args[2], env) end if expr.type == "ln" then return math.log(sym_eval(expr.args[1], env)) end if expr.type == "div" then return sym_eval(expr.args[1], env) / sym_eval(expr.args[2], env) end end -- 测试:求f(x)=x³+eˣ在x=1处的导数 local x_var = sym_exp("var", "x") local f = sym_exp("add", sym_exp("pow", x_var, 3), sym_exp("exp", x_var)) local f_prime = sym_diff(f, "x") print(sym_eval(f_prime, {x=1})) -- 精确值3+e≈5.71828
方案3:自动微分(前向模式)
利用链式法则,跟踪每个运算的导数,兼顾通用性和精度:
-- 定义对偶数值类型:存储值和导数 local function dual(val, der) return {val = val, der = der or 0} end -- 重载基础运算 local function dual_add(a, b) return dual(a.val + b.val, a.der + b.der) end local function dual_mul(a, b) return dual(a.val * b.val, a.der * b.val + a.val * b.der) end local function dual_pow(a, b) if type(b) == "number" then return dual(a.val ^ b, b * (a.val ^ (b-1)) * a.der) else local lnu = math.log(a.val) return dual(a.val ^ b.val, (b.der * lnu + b.val * a.der / a.val) * (a.val ^ b.val)) end end local function dual_exp(a) return dual(math.exp(a.val), math.exp(a.val) * a.der) end local function dual_sin(a) return dual(math.sin(a.val), math.cos(a.val) * a.der) end -- 自动微分求导函数 function derivative_auto(fun, x) local x_dual = dual(x, 1) -- 自变量的导数为1 local result = fun(x_dual) return result.der end -- 测试用例 print(derivative_auto(function(x) return dual_pow(x, 3) end, 10)) -- 300 print(derivative_auto(function(x) return dual_exp(x) end, 1)) -- ≈2.71828 print(derivative_auto(function(x) return dual_add(dual_pow(x, 2), dual_sin(x)) end, math.pi/2)) -- π≈3.14159
方案选择建议
- 若不需要极高精度,自适应数值微分最易用,支持任意可微函数
- 若需要完全精确的导数结果,且函数表达式已知,选择符号微分
- 若兼顾通用性和高精度,自动微分是最优选择,适合大多数场景
内容的提问来源于stack exchange,提问作者frogl8
相关产品推荐
相关产品推荐

