forward 教程
本教程用 diff 和 value_and_diff 对单变量函数求导,这是从本仓库获得导数的最简捷方式。你将对闭包、具名泛型函数以及调用带检查运算的函数求导,最后计算二阶导数。背景知识见 forward 设计。
快速上手
moon add Luna-Flow/autodiff@0.2.0
import {
"Luna-Flow/autodiff",
}
fn main {
let d = @autodiff.diff(x => x * x + x.sin(), 2.0)
println("d/dx (x^2 + sin x) at 2 = \{d}")
}
d/dx (x^2 + sin x) at 2 = 3.5838531634528574
闭包接收的是 Dual[Double],因此 * 和 sin 都是对偶运算,diff 返回 。
日常任务
同时获取值和导数
value_and_diff 通过一次求值同时返回两者:
fn main {
let (value, slope) = @autodiff.value_and_diff(x => x.exp() / (x * x), 2.0)
println("f(2) = \{value}")
println("f'(2) = \{slope}")
}
f(2) = 1.8472640247326626
f'(2) = 0
的导数为 ,它在 处为零。
在函数内部使用常数
捕获的数值必须以常数形式参与计算:
fn main {
let k = 3.0
let d = @autodiff.diff(
x => @autodiff.Dual::constant(k) * x * x - x.cos(),
1.0,
)
println("d/dx (3x^2 - cos x) at 1 = \{d}")
}
d/dx (3x^2 - cos x) at 1 = 6.841470984807897
对具名泛型函数求导
针对 trait 编写的函数可以直接传入;它会在 Dual[Double] 处实例化:
fn[T : @autodiff.Ring + @autodiff.Logarithmic] x_ln_x(x : T) -> T {
x * @autodiff.Logarithmic::ln(x)
}
fn main {
for x in [0.5, 1.0, 2.0] {
println("d/dx x ln x at \{x} = \{@autodiff.diff(x_ln_x, x)}")
}
}
d/dx x ln x at 0.5 = 0.3068528194400547
d/dx x ln x at 1 = 1
d/dx x ln x at 2 = 1.6931471805599454
导数为 。
用牛顿法求根
value_and_diff 恰好给出牛顿迭代一步所需的内容,即 :
fn main {
let f = (x : @autodiff.Dual[Double]) => x * x * x - @autodiff.Dual::constant(2.0)
let mut x = 1.0
for _ in 0..<5 {
let (fx, dfx) = @autodiff.value_and_diff(f, x)
x = x - fx / dfx
}
println("cube root of 2 = \{x}")
}
cube root of 2 = 1.2599210498948732
通过带检查的运算求导
当 f 使用 div_checked 或 sqrt_checked 时,请包装驱动函数的调用,使错误以值的形式离开闭包:
fn safe_derivative(x : Double) -> Result[Double, @autodiff.ArithmeticError] {
let ctx = @autodiff.ArithmeticContext::new(53)
let mut failure = None
let d = @autodiff.diff(
v => match v.sqrt_checked(ctx) {
Ok(r) => r
Err(e) => {
failure = Some(e)
@autodiff.Dual::constant(0.0)
}
},
x,
)
match failure {
Some(e) => Err(e)
None => Ok(d)
}
}
fn main {
for x in [4.0, -4.0] {
match safe_derivative(x) {
Ok(d) => println("sqrt'(\{x}) = \{d}")
Err(e) => println("sqrt'(\{x}) failed: domain error \{e.is_domain_error()}")
}
}
}
sqrt'(4) = 0.25
sqrt'(-4) failed: domain error true
深入学习
二阶导数
将 diff 应用于一个本身调用了 diff 的函数。内层函数必须是泛型的,才能在 Dual[Dual[Double]] 上运行:
fn[T : @autodiff.Ring + @autodiff.Trigonometric] f(x : T) -> T {
x * @autodiff.Trigonometric::sin(x)
}
fn main {
let second = @autodiff.diff(x => @autodiff.diff(f, x), 1.0)
let expected = 2.0 * @math.cos(1.0) - @math.sin(1.0)
println("f''(1) = \{second}")
println("2 cos 1 - sin 1 = \{expected}")
}
f''(1) = 0.23913362692838303
2 cos 1 - sin 1 = 0.23913362692838303
嵌套求得的导数与闭式解一致。每多一层工作量就翻倍,因此只适用于低阶导数。
其他标量类型
diff 适用于任何实现了 One 的 T;当 f 只使用环运算时,Float 甚至 Int 都可以:
fn main {
let d : Int = @autodiff.diff(x => x * x * x * x, 3)
println("d/dx x^4 at 3 = \{d}")
}
d/dx x^4 at 3 = 108
常见陷阱
- 闭包参数是对偶数。 请写
x.sin()或@autodiff.Trigonometric::sin(x),而不是@math.sin(x),后者接受Double,在这里无法编译。 - 捕获的值需要
Dual::constant。 如果捕获的Dual本身是Dual::variable(…),它会加入自己的切向分量,悄无声息地改变结果。 - 分支。
if x.value() < 0.0 { … }只对所走的分支求导。在拐点处你得到的是单侧导数。 - 嵌套需要泛型的内层函数。 嵌套
diff要求内层函数对T泛型;作用于Dual[Double]的闭包不能应用于Dual[Dual[Double]]。
后续阅读
- forward API 给出了驱动函数的确切约定。
- dual 教程 手动设定切向分量的种子,包括方向导数。
- linalg 教程 计算梯度和雅可比矩阵。
- forward 设计 解释了该方法以及嵌套的代价。