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
在原点处,平方根在 处求值,而它在那里没有导数,因此链条在此停止。
泛型的带检查代码
用带检查的 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在零处失败。 那里的导数不存在,即使输入是常数也是如此。- 极小的除数。 当 时 会下溢,切向分量的除法会报告除以零。
- 不带检查的运算符仍然不带检查。 对偶数上的
x / y和x.sqrt()从不返回错误;请使用_checked形式。
后续阅读
- checked API 列出了重新导出的名称。
- dual API 给出了精确的错误表。
- arithmetic 记录了错误模型。