You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求精准导数计算方案:现有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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.02 03:25:04