在 Julia 中绘制 ForwardDiff 的输出

Con*_*nor 3 type-conversion automatic-differentiation julia autodiff

我只想使用该ForwardDiff.jl功能来定义一个函数并绘制其梯度(使用 评估ForwardDiff.gradient)。它似乎不起作用,因为 的输出ForwardDiff.gradient是这种奇怪的Dual类型,并且不容易转换为所需的类型(在我的情况下,是 Float32s 的一维数组)。

using Plots
using ForwardDiff

my_func(x::Array{Float32,1}) = 1f0. / (1f0 .+ exp(3f0 .* x)) # doesn't matter what this is, just a sigmoid function here

grad_f(x::Array{Float32,1}) = ForwardDiff.gradient(my_func, x)

x_values = collect(Float32,0:0.01:10)

plot(x_values,my_func(x_values)); # this works fine

plot!(x_values,grad_f(x_values)); # this throws an error
Run Code Online (Sandbox Code Playgroud)

这是我得到的错误:

ERROR: MethodError: no method matching Float64(::ForwardDiff.Dual{ForwardDiff.Tag{typeof(g),Float32},Float64,12})
Run Code Online (Sandbox Code Playgroud)

当我检查 的类型时grad_f(x_values),我得到了这个:

Array{Array{ForwardDiff.Dual{ForwardDiff.Tag{typeof(g),Float32},Float32,12},1},1}

例如,为什么在 ForwardDiff 文档的示例中不会发生这种情况?见这里:https : //github.com/JuliaDiff/ForwardDiff.jl

提前致谢。

编辑:在 Kristoffer Carlsson 发表评论后:我试过这个,但它仍然不起作用。我不明白我在这里尝试的与他建议的有什么不同:

function g(x::Float32)
    return x / (1f0 + exp(10f0 * (x - 5f0)))
end

function ?g?x(x::Float32)
    return ForwardDiff.derivative(g, x)
end

x_vals = collect(Float32,0:0.01:10)
plot(x_vals,g.(x_vals))
plot!(x_vals,?g?x.(x_vals))
Run Code Online (Sandbox Code Playgroud)

现在的错误是:

no method matching g(::ForwardDiff.Dual{ForwardDiff.Tag{typeof(g),Float32},Float32,1})
Run Code Online (Sandbox Code Playgroud)

并且此错误仅在我调用时发生?g?x(x),无论我是否使用广播版本?g?x.(x)。我想这与函数定义有关,但我看不出我定义它的方式与 Kristoffer 的版本有何不同,除了它没有在一行中定义......这太令人困惑了。

这应该有效,因为根据ForwardDiff的文档,您只需要输入的类型是Real-的子类型,并且Float32是 Real 的子类型。

编辑:我现在意识到,在阅读了其他人的评论后,您需要将函数限制为足够通用以接受抽象类型的任何输入Real,而我并没有从文档中完全收集到这些输入。对混乱表示歉意。

Kri*_*son 5

您在数组而不是标量上定义函数,并且还过多地限制了输入类型。此外,对于标量函数,您应该使用ForwardDiff.derivative. 尝试类似:

using Plots
using ForwardDiff

my_func(x::Real) =  1f0 / (1f0 + exp(3f0 * x))
my_func_derivative(x::Real) = ForwardDiff.derivative(my_func, x)

plot(my_func, xlimits = (0, 10))
plot!(my_func_derivative)
Run Code Online (Sandbox Code Playgroud)

给予:

在此处输入图片说明

  • 您可能需要阅读 ForwardDiff 的工作原理 - 它使用[双数](https://en.wikipedia.org/wiki/Dual_number)调用您的函数来跟踪导数。如果您将函数限制为“Float32”类型的输入,则这将不起作用。您需要将其放宽为“Real”,因为 ForwardDiff 的“Dual”类型是“Real”的子类型。 (2认同)
  • 我想补充一点,除非您计划为该函数定义另一种方法(与其他类型的行为不同),否则您可能不应该强制执行任何输入类型。该代码将更易于阅读,同样具有高性能,并且很可能与其他包轻松地结合使用。 (2认同)