dual 教程
本教程用 Dual[T] 手动计算导数:为输入设定种子,执行普通的算术运算,再从切向分量中读出导数。学完之后,你将能够沿任意方向对多变量表达式求导,为普通数和对偶数编写同一个泛型函数,并把除法和平方根的失败作为数据处理。相关数学见 dual 设计。
快速上手
将该模块添加到你的项目中:
moon add Luna-Flow/autodiff@0.2.0
在你的 moon.pkg 中导入根包;它重新导出了 Dual 以及下文用到的 trait:
import {
"Luna-Flow/autodiff",
}
将 设为变量,并计算 :
fn main {
let x = @autodiff.Dual::variable(1.5)
let two = @autodiff.Dual::constant(2.0)
let y = x * x * x + two * x
println("f(1.5) = \{y.value()}")
println("f'(1.5) = \{y.tangent()}")
}
f(1.5) = 6.375
f'(1.5) = 8.75
切向分量为 。
日常任务
在计算中混入常数
所有不是求导变量的值都以切向分量为零的常数形式参与计算。字面量不能直接与对偶数混用,因此要用 Dual::constant 包装:
fn main {
let rate = 0.25
let x = @autodiff.Dual::variable(2.0)
let y = @autodiff.Dual::constant(rate) * x.exp()
println("d/dx 0.25 e^x at 2 = \{y.tangent()}")
}
d/dx 0.25 e^x at 2 = 1.8472640247326626
使用初等函数
sqrt、exp、exp2、ln、log2、log10、sin、cos 和 tan 都是 Dual[T] 的方法,会自动为你应用链式法则:
fn main {
let x = @autodiff.Dual::variable(1.0)
let y = x.sin().exp() // e^(sin x)
println("value = \{y.value()}")
println("derivative = \{y.tangent()}")
println("cos(1) e^(sin 1) = \{@math.cos(1.0) * @math.exp(@math.sin(1.0))}")
}
value = 2.319776824715853
derivative = 1.253380767493447
cos(1) e^(sin 1) = 1.253380767493447
第二行和第三行一致:切向分量为 。
沿选定方向求导
有多个输入时,你设定的切向分量决定了方向。以 、切向分量 设定种子,得到 ;切向分量取 ,则得到方向导数 :
fn f(x : @autodiff.Dual[Double], y : @autodiff.Dual[Double]) -> @autodiff.Dual[Double] {
x * y + x.sin()
}
fn main {
let dx = f(@autodiff.Dual::new(2.0, 1.0), @autodiff.Dual::new(3.0, 0.0))
let dy = f(@autodiff.Dual::new(2.0, 0.0), @autodiff.Dual::new(3.0, 1.0))
let dv = f(@autodiff.Dual::new(2.0, 0.5), @autodiff.Dual::new(3.0, -1.0))
println("df/dx = \{dx.tangent()}")
println("df/dy = \{dy.tangent()}")
println("grad . (0.5, -1) = \{dv.tangent()}")
}
df/dx = 2.5838531634528574
df/dy = 2
grad . (0.5, -1) = -0.7080734182735712
对于完整的梯度和雅可比矩阵,linalg 教程 中的驱动函数会替你完成这种种子设定。
与有限差分比较
中心差分需要选择步长,并会损失有效数字;对偶数的结果除舍入外是精确的:
fn g(x : Double) -> Double {
@math.exp(@math.sin(x))
}
fn main {
let exact = @math.cos(1.0) * g(1.0)
let dual = @autodiff.Dual::variable(1.0).sin().exp().tangent()
for h in [1.0e-2, 1.0e-5, 1.0e-8] {
let central = (g(1.0 + h) - g(1.0 - h)) / (2.0 * h)
println("h = \{h}: error \{(central - exact).abs()}")
}
println("dual: error \{(dual - exact).abs()}")
}
h = 0.01: error 0.00006752362465478612
h = 0.00001: error 7.4792172455318e-11
h = 1e-8: error 9.450921378828525e-9
dual: error 0
有限差分的误差先随 减小而下降,随后因抵消误差占主导而再次上升,与 dual 设计 中的推导完全一致。
为普通数和对偶数编写同一个函数
针对函数所需的 trait 编写它。这样同一份代码就能在 Double 上运行,求导时也能在 Dual[Double] 上运行。整数常量来自 IntegralHomomorphism::from_integral:
fn[T : @autodiff.Ring + @autodiff.IntegralHomomorphism + @autodiff.Trigonometric] h(
x : T,
) -> T {
let three : T = @autodiff.IntegralHomomorphism::from_integral(3)
three * x * x + @autodiff.Trigonometric::cos(x)
}
fn main {
println("h(0.5) = \{h(0.5)}")
let d = h(@autodiff.Dual::variable(0.5))
println("h(0.5) = \{d.value()} (dual value)")
println("h'(0.5) = \{d.tangent()}")
}
h(0.5) = 1.6275825618903728
h(0.5) = 1.6275825618903728 (dual value)
h'(0.5) = 2.520574461395797
在对偶数上计算出的值与在 Double 上计算出的值逐位相同;切向分量是 在 处的值。
报告除法与平方根的失败
div_checked 和 sqrt_checked 返回带有 arithmetic 错误的 Result,而不是无穷大或 NaN:
fn main {
let ctx = @autodiff.ArithmeticContext::new(53)
for a in [4.0, 0.0, -1.0] {
match @autodiff.Dual::variable(a).sqrt_checked(ctx) {
Ok(r) => println("sqrt(\{a}): value \{r.value()}, derivative \{r.tangent()}")
Err(e) =>
println(
"sqrt(\{a}): domain error \{e.is_domain_error()}, division by zero \{e.is_division_by_zero()}",
)
}
}
}
sqrt(4): value 2, derivative 0.25
sqrt(0): domain error false, division by zero true
sqrt(-1): domain error true, division by zero false
在 处平方根存在,但其导数不存在,因此切向分量的除法失败。
深入学习
通过嵌套求二阶导数
Dual[T] 是泛型的,因此 T 本身也可以是对偶数。对导数再求导要求内层函数是泛型的,forward 教程 用 diff 演示了这一点。手动实现如下:
fn[T : @autodiff.Ring] cube(x : T) -> T {
x * x * x
}
fn main {
// outer variable: tangent 1 on the outer level
let outer : @autodiff.Dual[Double] = @autodiff.Dual::variable(2.0)
// inner variable: the outer number, seeded with tangent 1 on the inner level
let inner = @autodiff.Dual::new(outer, @autodiff.Dual::constant(1.0))
let y = cube(inner)
println("f = \{y.value().value()}")
println("f' = \{y.tangent().value()}")
println("f'' = \{y.tangent().tangent()}")
}
f = 8
f' = 12
f'' = 12
两个层次是不同的类型,因此内层导数和外层导数的切向分量不会混淆。
你自己的标量类型
任何满足所调用方法约束的 T 都可以使用。像 Int 这样只有环结构的类型已经支持乘积法则:
fn main {
let x : @autodiff.Dual[Int] = @autodiff.Dual::variable(5)
let y = x * x * x
println("d/dx x^3 at 5 = \{y.tangent()}")
}
d/dx x^3 at 5 = 75
要让你自己的类型能通过 exp 或 sin 求导,请为它实现 Luna-Flow/arithmetic 中的 Exponential 或 Trigonometric。
常见陷阱
- 字面量不是对偶数。 当
x是Dual[Double]时,x * 2.0无法编译。请写成x * @autodiff.Dual::constant(2.0),或在泛型代码中使用IntegralHomomorphism::from_integral(2)。 - 比较只看到你所比较的东西。
Dual[T]没有<。请比较x.value();这样得到的导数就是所走分支的导数,因此用分支实现的abs之类的函数在拐点处得到的是单侧导数。 - 相等性包括切向分量。
Dual::new(1.0, 0.0) == Dual::new(1.0, 1.0)为false。 - 不带检查的运算沿用
Double的行为。 当y.value() == 0.0时的x / y、负数的ln或零处的sqrt都会在切向分量中产生无穷大或 NaN。若必须避免这种情况,请使用带检查的形式。 - 极小的除数。 商法则要除以 ,而当 时它会下溢;请在相除之前重新缩放。
- 无法直接打印。
Dual[T]没有Show。请打印value()和tangent(),或使用Debug提供的debug_inspect和@debug.to_string。
后续阅读
- dual API 列出了每个方法、规则和实例。
- dual 设计 推导了这些规则和误差界。
- forward 教程 把种子设定封装进
diff和value_and_diff;linalg 教程 计算梯度和雅可比矩阵;poly 教程 对多项式求导。 - 标量 trait 来自 luna-generic 和 arithmetic。