dual API
dual 包拥有 Dual[T],即满足 的对偶数 ,以及它的算术运算、Luna Flow trait 实例和初等函数的求导规则。本仓库的其他所有包都会重新导出或使用这个类型。这些规则背后的数学见 dual 设计。
源码:src/dual/dual.mbt 与 src/dual/extends.mbt。
导入
import {
"Luna-Flow/autodiff/dual",
}
大多数程序改为导入根包 Luna-Flow/autodiff,它将 Dual 重新导出为 @autodiff.Dual(见 autodiff API)。本页示例使用根包别名。
类型
Dual
Dual[T] 存储一个原始值和一个一阶切向分量。
pub struct Dual[T] {
value : T
tangent : T
} derive(Eq, @debug.Debug)
二元组 (value, tangent) 表示 ,其中 = value, = tangent。当计算从 Dual::variable(x) 开始时,每个中间结果的切向分量就是该中间量对 x 的导数。这些字段在包外是只读的;请用下面的构造函数构建值。
derive(Eq) 会比较两个分量,因此值相同而切向分量不同的两个对偶数并不相等。derive(Debug) 打印记录形式 { value: …, tangent: … },debug_inspect 和 assert_eq 都使用这种形式。
Dual[T] 有意不提供 Compare、Field、MulGroup、Inverse 或 Show 实例;见设计决策。
构造与访问
Dual::new
由显式给出的值和切向分量构建 。
pub fn[T] Dual::new(T, T) -> Dual[T]
用它来设定 以外的种子方向,例如计算方向导数 ,此时每个输入坐标的切向分量取 。
Dual::constant
以零切向分量嵌入一个值:。
pub fn[T : @luna-generic.Zero] Dual::constant(T) -> Dual[T]
常数不依赖于求导变量,因此其导数为 。凡是进入被求导计算的字面量或捕获值,都应使用它。
Dual::variable
以切向分量 1 嵌入求导变量:。
pub fn[T : @luna-generic.One] Dual::variable(T) -> Dual[T]
Dual::value
返回原始分量 。
pub fn[T] Dual::value(Dual[T]) -> T
Dual::tangent
返回切向分量 。
pub fn[T] Dual::tangent(Dual[T]) -> T
test "construct and read dual numbers" {
let x = @autodiff.Dual::new(2.0, 3.0)
assert_eq(x.value(), 2.0)
assert_eq(x.tangent(), 3.0)
let c : @autodiff.Dual[Double] = @autodiff.Dual::constant(5.0)
assert_eq(c.tangent(), 0.0)
let v : @autodiff.Dual[Double] = @autodiff.Dual::variable(5.0)
assert_eq(v.tangent(), 1.0)
}
Dual::zero
返回加法单位元 。
pub fn[T : @luna-generic.Zero] Dual::zero() -> Dual[T]
这是 Zero 实例提升而来的方法。
Dual::one
返回乘法单位元 。
pub fn[T : @luna-generic.One + @luna-generic.Zero] Dual::one() -> Dual[T]
one 是常数,所以其切向分量为零;它不是 Dual::variable(1)。这是 One 实例提升而来的方法。
算术运算
这四个运算符实现了和、差、积、商的求导法则。它们既可以作为运算符使用,也可以作为提升方法使用。
| 条目 | 运算符 | 规则 | 对 T 的约束 |
|---|---|---|---|
Dual::add | x + y | Add | |
Dual::sub | x - y | Sub | |
Dual::neg | -x | Neg | |
Dual::mul | x * y | Add + Mul | |
Dual::div | x / y | Div + Sub + Mul |
Dual::add
逐分量相加。
pub fn[T : Add] Dual::add(Dual[T], Dual[T]) -> Dual[T]
pub impl[T : Add] Add for Dual[T]
Dual::sub
逐分量相减。
pub fn[T : Sub] Dual::sub(Dual[T], Dual[T]) -> Dual[T]
pub impl[T : Sub] Sub for Dual[T]
Dual::neg
对两个分量取负。
pub fn[T : Neg] Dual::neg(Dual[T]) -> Dual[T]
pub impl[T : Neg] Neg for Dual[T]
Dual::mul
按乘积法则相乘。
pub fn[T : Add + Mul] Dual::mul(Dual[T], Dual[T]) -> Dual[T]
pub impl[T : Add + Mul] Mul for Dual[T]
切向分量按 self.value * other.tangent + other.value * self.tangent 计算,即 。由于所有数值类型 T 的乘法都满足交换律,它等于乘积法则 。
Dual::div
按商法则相除,不做任何定义域检查。
pub fn[T : Div + Sub + Mul] Dual::div(Dual[T], Dual[T]) -> Dual[T]
pub impl[T : Div + Sub + Mul] Div for Dual[T]
切向分量为 (self.tangent * other.value - self.value * other.tangent) / (other.value * other.value)。除数的值为零时,结果就是 T 的除法给出的结果(对于 Double,是无穷大或 NaN)。如果希望得到错误,请使用 Dual::div_checked。
test "dual arithmetic" {
let x = @autodiff.Dual::new(2.0, 3.0)
let y = @autodiff.Dual::new(5.0, 7.0)
assert_eq(x + y, @autodiff.Dual::new(7.0, 10.0))
assert_eq(y - x, @autodiff.Dual::new(3.0, 4.0))
assert_eq(-x, @autodiff.Dual::new(-2.0, -3.0))
assert_eq(x * y, @autodiff.Dual::new(10.0, 29.0)) // 2*7 + 5*3
assert_eq(y / x, @autodiff.Dual::new(2.5, -0.25)) // (7*2 - 5*3) / 4
}
Dual::equal
判断两个分量是否都相等。
pub fn[T : Eq] Dual::equal(Dual[T], Dual[T]) -> Bool
这是派生的 Eq 实例提升而来的方法;推荐使用 ==。
带检查的运算
带检查的形式返回来自 Luna-Flow/arithmetic 的 Result[Dual[T], ArithmeticError],而不会产生无穷大或 NaN。checked API 重新导出了它们使用的类型。
Dual::div_checked
按商法则相除,并将非法除法报告为错误。
pub fn[T : @arithmetic.DivChecked + Sub + Mul] Dual::div_checked(Dual[T], Dual[T], @arithmetic.ArithmeticContext) -> Result[Dual[T], @arithmetic.ArithmeticError]
pub impl[T : @arithmetic.DivChecked + Sub + Mul] @arithmetic.DivChecked for Dual[T]
先用 T 的 div_checked 计算值 ,再用第二次 div_checked 调用计算切向分量 。遇到的第一个错误会原样返回。对于 Double,错误如下:
| 条件 | 错误 |
|---|---|
| 且 | is_division_by_zero() |
| 且 | is_domain_error()(零除以零) |
| 与 均为无穷大 | is_domain_error() |
| 但 下溢为 | 切向分量的除法失败;见下方的陷阱说明 |
上下文参数会被传递给 T;arithmetic 的 Double 与 Float 实例会忽略它。
test "checked dual division" {
let ctx = @autodiff.ArithmeticContext::new(53)
let x = @autodiff.Dual::new(6.0, 2.0)
let y = @autodiff.Dual::new(3.0, 1.0)
match x.div_checked(y, ctx) {
Ok(q) => assert_eq(q, @autodiff.Dual::new(2.0, 0.0))
Err(_) => fail("unexpected error")
}
match x.div_checked(@autodiff.Dual::new(0.0, 1.0), ctx) {
Ok(_) => fail("expected an error")
Err(e) => assert_true(e.is_division_by_zero())
}
}
Dual::sqrt_checked
按法则 求平方根,并报告定义域错误。
pub fn[T : @arithmetic.SqrtChecked + @arithmetic.DivChecked + @luna-generic.IntegralHomomorphism + Mul] Dual::sqrt_checked(Dual[T], @arithmetic.ArithmeticContext) -> Result[Dual[T], @arithmetic.ArithmeticError]
pub impl[T : @arithmetic.SqrtChecked + @arithmetic.DivChecked + @luna-generic.IntegralHomomorphism + Mul] @arithmetic.SqrtChecked for Dual[T]
值为 T 的 sqrt_checked(a);切向分量为 div_checked(b, 2 * root),其中 2 来自 IntegralHomomorphism::from_integral(2)。对于 Double:
| 条件 | 错误 |
|---|---|
来自平方根的 is_domain_error() | |
| , | 来自切向分量的 is_division_by_zero() |
| , | 来自切向分量的 is_domain_error()(零除以零) |
因此即使输入是常数,sqrt_checked 在 处也会失败: 在 处没有导数,而带检查的形式不会对零切向分量做特殊处理。
test "checked dual square root" {
let ctx = @autodiff.ArithmeticContext::new(53)
match @autodiff.Dual::new(9.0, 6.0).sqrt_checked(ctx) {
Ok(r) => assert_eq(r, @autodiff.Dual::new(3.0, 1.0))
Err(_) => fail("unexpected error")
}
match @autodiff.Dual::new(-1.0, 1.0).sqrt_checked(ctx) {
Ok(_) => fail("expected an error")
Err(e) => assert_true(e.is_domain_error())
}
}
初等函数
每个初等函数既是固有方法,也是对应 arithmetic trait 实例的方法,因此针对 Sqrt、Exponential、Logarithmic 或 Trigonometric 编写的泛型代码无需修改即可通过 Dual[T] 求导。它们都使用下表中的导数,应用法则 。它们都不检查定义域:超出定义域时,结果就是 T 返回的结果(对于 Double,是 NaN 或无穷大)。
| 条目 | 值 | 切向分量的实际计算方式 | 导数 |
|---|---|---|---|
Dual::sqrt | b / (2 * sqrt(a)) | ||
Dual::exp | b * exp(a) | ||
Dual::exp2 | b * exp2(a) * ln(2) | ||
Dual::ln | b / a | ||
Dual::log2 | b / (a * ln(2)) | ||
Dual::log10 | b / (a * ln(10)) | ||
Dual::sin | b * cos(a) | ||
Dual::cos | -(b * sin(a)) | ||
Dual::tan | b / (cos(a) * cos(a)) |
常数 和 来自 IntegralHomomorphism::from_integral。
Dual::sqrt
平方根,切向分量为 。
pub fn[T : @arithmetic.Sqrt + @luna-generic.IntegralHomomorphism + Mul + Div] Dual::sqrt(Dual[T]) -> Dual[T]
pub impl[T : @arithmetic.Sqrt + @luna-generic.IntegralHomomorphism + Mul + Div] @arithmetic.Sqrt for Dual[T]
Dual::exp
指数函数,切向分量为 ;值只计算一次并被复用。
pub fn[T : @arithmetic.Exponential + Mul] Dual::exp(Dual[T]) -> Dual[T]
Dual::exp2
以 2 为底的指数函数,切向分量为 。
pub fn[T : @arithmetic.Exponential + @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul] Dual::exp2(Dual[T]) -> Dual[T]
pub impl[T : @arithmetic.Exponential + @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul] @arithmetic.Exponential for Dual[T]
Exponential 实例同时提供 exp 和 exp2,因此需要 exp2 的约束。
Dual::ln
自然对数,切向分量为 。
pub fn[T : @arithmetic.Logarithmic + Div] Dual::ln(Dual[T]) -> Dual[T]
Dual::log2
以 2 为底的对数,切向分量为 。
pub fn[T : @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul + Div] Dual::log2(Dual[T]) -> Dual[T]
Dual::log10
以 10 为底的对数,切向分量为 。
pub fn[T : @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul + Div] Dual::log10(Dual[T]) -> Dual[T]
pub impl[T : @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul + Div] @arithmetic.Logarithmic for Dual[T]
Logarithmic 实例提供 ln、log2 和 log10。
Dual::sin
正弦,切向分量为 。
pub fn[T : @arithmetic.Trigonometric + Mul] Dual::sin(Dual[T]) -> Dual[T]
Dual::cos
余弦,切向分量为 。
pub fn[T : @arithmetic.Trigonometric + Mul + Neg] Dual::cos(Dual[T]) -> Dual[T]
Dual::tan
正切,导数为 ,按 计算。
pub fn[T : @arithmetic.Trigonometric + Mul + Div] Dual::tan(Dual[T]) -> Dual[T]
pub impl[T : @arithmetic.Trigonometric + Mul + Neg + Div] @arithmetic.Trigonometric for Dual[T]
Trigonometric 实例提供 sin、cos 和 tan。
test "elementary derivative rules" {
let x = @autodiff.Dual::variable(0.5)
assert_eq(x.sin().tangent(), @math.cos(0.5))
assert_eq(x.exp().tangent(), @math.exp(0.5))
assert_eq(x.ln().tangent(), 2.0)
assert_eq(@autodiff.Dual::variable(4.0).sqrt().tangent(), 0.25)
}
Trait 实例
Dual[T] 实现了代数 所满足的 Luna Flow 结构 trait,每个都要求 T 满足相应的约束。这些定律的推导见 dual 设计。
| 实例 | 对 T 的约束 | 含义 |
|---|---|---|
Zero | Zero | zero() 为 |
One | One + Zero | one() 为 |
AddMonoid | AddMonoid | 逐分量加法 |
AddGroup | AddGroup | 逐分量取负 |
MulMonoid | Semiring | 按乘积法则的乘法 |
Semiring | Semiring | 当 T 是半环时 是半环 |
Ring | Ring | 当 T 是环时 是环 |
NatHomomorphism | NatHomomorphism + Zero | from_nat(n) 即 Dual::constant(from_nat(n)) |
IntegralHomomorphism | IntegralHomomorphism + Zero | from_integral(n) 即 Dual::constant(from_integral(n)) |
@arithmetic.Constants | Constants + Zero | pi()、tau()、e() 都是常数 |
@arithmetic.DivChecked | DivChecked + Sub + Mul | 见 Dual::div_checked |
@arithmetic.SqrtChecked | 见 Dual::sqrt_checked | 见 Dual::sqrt_checked |
@arithmetic.Sqrt, Exponential, Logarithmic, Trigonometric | 见各个方法 | 见初等函数 |
整数与自然数的嵌入以及常数的切向分量都为零,因为它们不依赖于求导变量。
fn[T : @autodiff.Ring + @autodiff.IntegralHomomorphism] three_x_squared(x : T) -> T {
let three : T = @autodiff.IntegralHomomorphism::from_integral(3)
three * x * x
}
test "generic code differentiates through the trait instances" {
let y = three_x_squared(@autodiff.Dual::variable(2.0))
assert_eq(y, @autodiff.Dual::new(12.0, 12.0))
}
已弃用
Dual 保留了几个方法形式,它们过去由 MoonBit 根据 trait 实例隐式生成。这些方法不会出现在接口文件中,在本包之外使用时会产生警告。
| 已弃用的方法 | 替代方式 |
|---|---|
x.not_equal(y) | x != y |
x.to_repr() | Repr(x) 或 @debug.Debug::to_repr(x) |
Dual::from_nat(n) | 来自 Luna-Flow/luna-generic 的 NatHomomorphism::from_nat(n) |
Dual::from_integral(n) | @autodiff.IntegralHomomorphism::from_integral(n) |
Dual::pi(), Dual::e(), Dual::tau() | @autodiff.Constants::pi(), e(), tau() |