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 返回 2⋅2+cos⁡22 \cdot 2 + \cos 2。

日常任务

同时获取值和导数

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

f(x)=ex/x2f(x) = e^x / x^2 的导数为 f′(x)=ex(x−2)/x3f'(x) = e^x (x - 2)/x^3,它在 x=2x = 2 处为零。

在函数内部使用常数

捕获的数值必须以常数形式参与计算:

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

导数为 ln⁡x+1\ln x + 1。

用牛顿法求根

value_and_diff 恰好给出牛顿迭代一步所需的内容,即 x←x−f(x)/f′(x)x \leftarrow x - f(x)/f'(x):

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]]。

后续阅读