dual API

dual 包拥有 Dual[T],即满足 ε2=0\varepsilon^2 = 0 的对偶数 a+bεa + b\varepsilon,以及它的算术运算、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) 表示 a+bεa + b\varepsilon,其中 aa = value,bb = tangent。当计算从 Dual::variable(x) 开始时,每个中间结果的切向分量就是该中间量对 x 的导数。这些字段在包外是只读的;请用下面的构造函数构建值。

derive(Eq) 会比较两个分量,因此值相同而切向分量不同的两个对偶数并不相等。derive(Debug) 打印记录形式 { value: …, tangent: … },debug_inspect 和 assert_eq 都使用这种形式。

Dual[T] 有意不提供 Compare、Field、MulGroup、Inverse 或 Show 实例;见设计决策。

构造与访问

Dual::new

由显式给出的值和切向分量构建 a+bεa + b\varepsilon。

pub fn[T] Dual::new(T, T) -> Dual[T]

用它来设定 11 以外的种子方向,例如计算方向导数 ∇f(x)⋅v\nabla f(x) \cdot v,此时每个输入坐标的切向分量取 viv_i。

Dual::constant

以零切向分量嵌入一个值:c↦c+0εc \mapsto c + 0\varepsilon。

pub fn[T : @luna-generic.Zero] Dual::constant(T) -> Dual[T]

常数不依赖于求导变量,因此其导数为 00。凡是进入被求导计算的字面量或捕获值,都应使用它。

Dual::variable

以切向分量 1 嵌入求导变量:x↦x+1εx \mapsto x + 1\varepsilon。

pub fn[T : @luna-generic.One] Dual::variable(T) -> Dual[T]

Dual::value

返回原始分量 aa。

pub fn[T] Dual::value(Dual[T]) -> T

Dual::tangent

返回切向分量 bb。

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

返回加法单位元 0+0ε0 + 0\varepsilon。

pub fn[T : @luna-generic.Zero] Dual::zero() -> Dual[T]

这是 Zero 实例提升而来的方法。

Dual::one

返回乘法单位元 1+0ε1 + 0\varepsilon。

pub fn[T : @luna-generic.One + @luna-generic.Zero] Dual::one() -> Dual[T]

one 是常数,所以其切向分量为零;它不是 Dual::variable(1)。这是 One 实例提升而来的方法。

算术运算

这四个运算符实现了和、差、积、商的求导法则。它们既可以作为运算符使用,也可以作为提升方法使用。

条目运算符规则对 T 的约束
Dual::addx + y(a+bε)+(c+dε)=(a+c)+(b+d)ε(a+b\varepsilon)+(c+d\varepsilon) = (a+c) + (b+d)\varepsilonAdd
Dual::subx - y(a−c)+(b−d)ε(a-c) + (b-d)\varepsilonSub
Dual::neg-x−a−bε-a - b\varepsilonNeg
Dual::mulx * yac+(ad+cb)εac + (ad + cb)\varepsilonAdd + Mul
Dual::divx / yac+bc−adc2ε\dfrac{a}{c} + \dfrac{bc - ad}{c^2}\varepsilonDiv + 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 计算,即 ad+cbad + cb。由于所有数值类型 T 的乘法都满足交换律,它等于乘积法则 ad+bcad + bc。

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 计算值 a/ca / c,再用第二次 div_checked 调用计算切向分量 (bc−ad)/c2(bc - ad) / c^2。遇到的第一个错误会原样返回。对于 Double,错误如下:

条件错误
c=0c = 0 且 a≠0a \ne 0is_division_by_zero()
a=0a = 0 且 c=0c = 0is_domain_error()(零除以零)
aa 与 cc 均为无穷大is_domain_error()
c≠0c \ne 0 但 c⋅cc \cdot c 下溢为 00切向分量的除法失败;见下方的陷阱说明

上下文参数会被传递给 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

按法则 a+bε=a+b2aε\sqrt{a + b\varepsilon} = \sqrt a + \dfrac{b}{2\sqrt a}\varepsilon 求平方根,并报告定义域错误。

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:

条件错误
a<0a < 0来自平方根的 is_domain_error()
a=0a = 0, b≠0b \ne 0来自切向分量的 is_division_by_zero()
a=0a = 0, b=0b = 0来自切向分量的 is_domain_error()(零除以零)

因此即使输入是常数,sqrt_checked 在 a=0a = 0 处也会失败:⋅\sqrt{\cdot} 在 00 处没有导数,而带检查的形式不会对零切向分量做特殊处理。

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] 求导。它们都使用下表中的导数,应用法则 f(a+bε)=f(a)+f′(a) b εf(a + b\varepsilon) = f(a) + f'(a)\,b\,\varepsilon。它们都不检查定义域:超出定义域时,结果就是 T 返回的结果(对于 Double,是 NaN 或无穷大)。

条目值切向分量的实际计算方式导数
Dual::sqrta\sqrt ab / (2 * sqrt(a))12a\frac{1}{2\sqrt a}
Dual::expeae^ab * exp(a)eae^a
Dual::exp22a2^ab * exp2(a) * ln(2)2aln⁡22^a \ln 2
Dual::lnln⁡a\ln ab / a1a\frac{1}{a}
Dual::log2log⁡2a\log_2 ab / (a * ln(2))1aln⁡2\frac{1}{a\ln 2}
Dual::log10log⁡10a\log_{10} ab / (a * ln(10))1aln⁡10\frac{1}{a\ln 10}
Dual::sinsin⁡a\sin ab * cos(a)cos⁡a\cos a
Dual::coscos⁡a\cos a-(b * sin(a))−sin⁡a-\sin a
Dual::tantan⁡a\tan ab / (cos(a) * cos(a))sec⁡2a\sec^2 a

常数 22 和 1010 来自 IntegralHomomorphism::from_integral。

Dual::sqrt

平方根,切向分量为 b/(2a)b / (2\sqrt a)。

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

指数函数,切向分量为 b eab\,e^a;值只计算一次并被复用。

pub fn[T : @arithmetic.Exponential + Mul] Dual::exp(Dual[T]) -> Dual[T]

Dual::exp2

以 2 为底的指数函数,切向分量为 b⋅2aln⁡2b \cdot 2^a \ln 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

自然对数,切向分量为 b/ab / a。

pub fn[T : @arithmetic.Logarithmic + Div] Dual::ln(Dual[T]) -> Dual[T]

Dual::log2

以 2 为底的对数,切向分量为 b/(aln⁡2)b / (a \ln 2)。

pub fn[T : @arithmetic.Logarithmic + @luna-generic.IntegralHomomorphism + Mul + Div] Dual::log2(Dual[T]) -> Dual[T]

Dual::log10

以 10 为底的对数,切向分量为 b/(aln⁡10)b / (a \ln 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

正弦,切向分量为 bcos⁡ab \cos a。

pub fn[T : @arithmetic.Trigonometric + Mul] Dual::sin(Dual[T]) -> Dual[T]

Dual::cos

余弦,切向分量为 −bsin⁡a-b \sin a。

pub fn[T : @arithmetic.Trigonometric + Mul + Neg] Dual::cos(Dual[T]) -> Dual[T]

Dual::tan

正切,导数为 sec⁡2a\sec^2 a,按 b/(cos⁡a⋅cos⁡a)b / (\cos a \cdot \cos a) 计算。

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] 实现了代数 T[ε]/(ε2)T[\varepsilon]/(\varepsilon^2) 所满足的 Luna Flow 结构 trait,每个都要求 T 满足相应的约束。这些定律的推导见 dual 设计。

实例对 T 的约束含义
ZeroZerozero() 为 0+0ε0 + 0\varepsilon
OneOne + Zeroone() 为 1+0ε1 + 0\varepsilon
AddMonoidAddMonoid逐分量加法
AddGroupAddGroup逐分量取负
MulMonoidSemiring按乘积法则的乘法
SemiringSemiring当 T 是半环时 T[ε]T[\varepsilon] 是半环
RingRing当 T 是环时 T[ε]T[\varepsilon] 是环
NatHomomorphismNatHomomorphism + Zerofrom_nat(n) 即 Dual::constant(from_nat(n))
IntegralHomomorphismIntegralHomomorphism + Zerofrom_integral(n) 即 Dual::constant(from_integral(n))
@arithmetic.ConstantsConstants + Zeropi()、tau()、e() 都是常数
@arithmetic.DivCheckedDivChecked + 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()