checked 教程

本教程对可能失败的代码求导。你将在对偶数上调用带检查的除法和平方根,在计算中传播它们的错误,并编写能同时在普通数和对偶数上运行的泛型带检查代码。背后的原理见 checked 设计。

快速上手

moon add Luna-Flow/autodiff@0.2.0
import {
  "Luna-Flow/autodiff/checked",
}
fn main {
  let ctx = @checked.ArithmeticContext::new(53)
  let x : @checked.Dual[Double] = @checked.Dual::variable(2.0)
  let y : @checked.Dual[Double] = @checked.Dual::constant(8.0)
  match y.div_checked(x, ctx) {
    Ok(q) => println("8/x at 2: value \{q.value()}, derivative \{q.tangent()}")
    Err(e) => println("failed: \{e.message}")
  }
}
8/x at 2: value 4, derivative -2

日常任务

检查失败信息

fn main {
  let ctx = @checked.ArithmeticContext::new(53)
  let one : @checked.Dual[Double] = @checked.Dual::constant(1.0)
  for c in [1.0, 0.0] {
    match one.div_checked(@checked.Dual::variable(c), ctx) {
      Ok(q) => println("1/x at \{c}: derivative \{q.tangent()}")
      Err(e) => println("1/x at \{c}: \{e.message}")
    }
  }
}
1/x at 1: derivative -1
1/x at 0: division by zero

串联带检查的步骤

每一步都返回一个 Result;在遇到第一个错误时停止:

fn norm_ratio(
  x : @checked.Dual[Double],
  y : @checked.Dual[Double],
  ctx : @checked.ArithmeticContext,
) -> Result[@checked.Dual[Double], @checked.ArithmeticError] {
  // sqrt(x^2 + y^2) / y
  let r = (x * x + y * y).sqrt_checked(ctx)
  match r {
    Err(e) => Err(e)
    Ok(r) => r.div_checked(y, ctx)
  }
}

fn main {
  let ctx = @checked.ArithmeticContext::new(53)
  let y : @checked.Dual[Double] = @checked.Dual::constant(4.0)
  match norm_ratio(@checked.Dual::variable(3.0), y, ctx) {
    Ok(v) => println("value \{v.value()}, d/dx \{v.tangent()}")
    Err(e) => println(e.message)
  }
  let zero : @checked.Dual[Double] = @checked.Dual::constant(0.0)
  match norm_ratio(@checked.Dual::variable(0.0), zero, ctx) {
    Ok(_) => println("unexpected")
    Err(e) => println("at the origin: \{e.message}")
  }
}
value 1.25, d/dx 0.15
at the origin: zero divided by zero is undefined

在原点处,平方根在 00 处求值,而它在那里没有导数,因此链条在此停止。

泛型的带检查代码

用带检查的 trait 约束函数;这样它就能在 Double 和 Dual[Double] 上运行:

fn[T : @checked.DivChecked + @checked.SqrtChecked] sqrt_ratio(
  a : T,
  b : T,
  ctx : @checked.ArithmeticContext,
) -> Result[T, @checked.ArithmeticError] {
  match @checked.DivChecked::div_checked(a, b, ctx) {
    Err(e) => Err(e)
    Ok(q) => @checked.SqrtChecked::sqrt_checked(q, ctx)
  }
}

fn main {
  let ctx = @checked.ArithmeticContext::new(53)
  match sqrt_ratio(8.0, 2.0, ctx) {
    Ok(v) => println("plain: \{v}")
    Err(e) => println(e.message)
  }
  let a : @checked.Dual[Double] = @checked.Dual::variable(8.0)
  let b : @checked.Dual[Double] = @checked.Dual::constant(2.0)
  match sqrt_ratio(a, b, ctx) {
    Ok(v) => println("dual: \{v.value()}, d/da \{v.tangent()}")
    Err(e) => println(e.message)
  }
}
plain: 2
dual: 2, d/da 0.125

深入学习

  • ArithmeticContext::new(precision, rounding=…) 为使用上下文的标量类型构建上下文;Double 和 Float 会忽略它。
  • 传给 @autodiff.diff 的函数可以使用带检查的运算;forward 教程 展示了如何把错误带出闭包。
  • 自定义标量类型只需实现 DivChecked 和 SqrtChecked 即可参与;其错误随后会原样出现在对偶数结果中。

常见陷阱

  • sqrt_checked 在零处失败。 那里的导数不存在,即使输入是常数也是如此。
  • 极小的除数。 当 ∣c∣<1.5×10−162|c| < 1.5 \times 10^{-162} 时 c2c^2 会下溢,切向分量的除法会报告除以零。
  • 不带检查的运算符仍然不带检查。 对偶数上的 x / y 和 x.sqrt() 从不返回错误;请使用 _checked 形式。

后续阅读