main.bend checks
raw source on the hub · import bend-ml-autograd@0.1.1.0/main.bend as Main
bend-ml-autograd: automatic differentiation with a proved law.
import bend-ml-autograd@0.1.1.0/main.bend as AG
1) PROVED MODEL (over Nat): the reverse mode of autodiff gives the same result as the forward mode (dual numbers). Proved by the kernel, using the sum and product lemmas of bend-ml-nat-lemmas. 2) SCALAR AUTOGRAD in F32 (micrograd style), with the same structure. 3) LAYERS WITH TENSORS (forward and backward) whose shapes are checked by the type.
F32 is not a real number, so the NUMERICAL correctness of F32 is validated by tests against PyTorch (reference/test_autograd.py), not by proof.
3 imports
import Base import bend-ml-nat-lemmas@0.1.0.0/main.bend as NL import bend-ml-tensor@0.1.0.0/main.bend as T
Laws
law reverse_eq_forward provedsource · line 113 · raw
@+e:NE -> @+x:Nat -> {nbwd(e, x, 1n) == nfwd(e, x) : Nat}LAW (reverse_eq_forward): for any expression made of constants, X, sums and products, the reverse-mode gradient (starting with gradient 1 at the output) equals the forward-mode derivative. It holds for every x.
Types
type NE source · line 23 · raw
Data
Expressions with one variable NX, over Nat (a commutative semiring).
NCst@n:Nat -> NE
NXNE
NAdd@a:NE -> @b:NE -> NE
NMul@a:NE -> @b:NE -> NE
type G source · line 126 · raw
Data
GCst@c:F32 -> G
GVar@i:Nat -> G
GAdd@a:G -> @b:G -> G
GMul@a:G -> @b:G -> G
GRelu@a:G -> G
GTanh@a:G -> G
GExp@a:G -> G
Definitions
def nval source · line 30 · raw
@e:NE -> @+x:Nat -> Nat
value of the expression at x
def nfwd source · line 42 · raw
@e:NE -> @+x:Nat -> Nat
FORWARD-mode derivative (dual numbers): derivative of the sum and the product rule
def nbwd source · line 56 · raw
@e:NE -> @+x:Nat -> @+g:Nat -> Nat
REVERSE-mode derivative: g is the gradient arriving from the output; each node passes on g (sum) or g * value-of-the-other-factor (product) to its children and adds up what arrives at NX.
def ndist_l source · line 68 · raw
@+a:Nat -> @+b:Nat -> @+c:Nat -> {Nat.mul(a, Nat.add(b, c)) == Nat.add(Nat.mul(a, b), Nat.mul(a, c)) : Nat}a * (b + c) == a*b + a*c (left distributivity; nat-lemmas only has the right one)
def nmul_case source · line 79 · raw
@+g:Nat -> @+fa:Nat -> @+fb:Nat -> @+va:Nat -> @+vb:Nat -> @+ba:Nat -> @+bb:Nat -> @iha:{ba == Nat.mul(Nat.mul(g, vb), fa) : Nat} -> @ihb:{bb == Nat.mul(Nat.mul(g, va), fb) : Nat} -> {Nat.add(ba, bb) == Nat.mul(g, Nat.add(Nat.mul(fa, vb), Nat.mul(va, fb))) : Nat}The product case, with everything abstracted into numbers: (g*vb)*fa + (g*va)*fb == g * (fa*vb + va*fb)
def nbwd_ok source · line 95 · raw
@e:NE -> @+x:Nat -> @+g:Nat -> {nbwd(e, x, g) == Nat.mul(g, nfwd(e, x)) : Nat}The reverse mode returns g times the forward-mode derivative.
def lookup source · line 135 · raw
@env:List<&2, F32> -> @+i:Nat -> F32
def gval source · line 144 · raw
@e:G -> @+env:List<&2, F32> -> F32
def acc_add source · line 162 · raw
@acc:List<&2, F32> -> @+i:Nat -> @+v:F32 -> List<&2, F32>
adds v to the gradient of variable i
def mask2 source · line 171 · raw
@pos:Bool -> F32
def relu_d source · line 179 · raw
@+x:F32 -> F32
local derivative of relu: 1 if the input > 0, else 0
def tanh_d source · line 182 · raw
@+x:F32 -> F32
def gbwd source · line 186 · raw
@e:G -> @+env:List<&2, F32> -> @+g:F32 -> @acc:List<&2, F32> -> List<&2, F32>
g is the gradient arriving at the node; acc accumulates the gradient of each variable
def ggrad source · line 204 · raw
@e:G -> @+env:List<&2, F32> -> List<&2, F32>
gradient of the expression with respect to each variable of env
def linear source · line 212 · raw
@-n:Nat -> @+i:Nat -> @+o:Nat -> @x:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, i> -> @w:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<i, o> -> @b:0xf9837737d2c3f58ae1d0c5df42255584/main.Vec<o> -> 0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, o>
y = x·W + b x: n x i, W: i x o, b: o -> y: n x o
def linear_bwd source · line 220 · raw
@+n:Nat -> @+i:Nat -> @+o:Nat -> @x:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, i> -> @w:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<i, o> -> @+dy:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, o> -> Pair(0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, i>, Pair(0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<i, o>, 0xf9837737d2c3f58ae1d0c5df42255584/main.Vec<o>))
Given the gradient dy arriving at y, returns (dx, (dW, db)): dx = dy · Wᵀ (n x o)·(o x i) = n x i dW = xᵀ · dy (i x n)·(n x o) = i x o db = sum of the rows of dy If any of these products had the wrong dimensions, the program would not compile.
def mask.l source · line 223 · raw
@xs:List<&2, F32> -> List<&2, F32>
def mask.rows source · line 230 · raw
@xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>
def relu_bwd source · line 238 · raw
@-r:Nat -> @-c:Nat -> @x:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c> -> @dy:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c> -> 0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c>
y = relu(x); the gradient passes where x > 0
def onehot.cell source · line 244 · raw
@hit:Bool -> F32
one-hot: n rows of c numbers, with 1 at the label's position
def onehot.row source · line 247 · raw
@c:Nat -> @+k:Nat -> @+j:Nat -> List<&2, F32>
def onehot.rows source · line 254 · raw
@labels:List<&2, Nat> -> @+c:Nat -> List<&2, List<&2, F32>>
def log.l source · line 261 · raw
@xs:List<&2, F32> -> List<&2, F32>
def log.rows source · line 268 · raw
@xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>
def rows.total source · line 275 · raw
@xs:List<&2, List<&2, F32>> -> @acc:F32 -> F32
def ce_loss2 source · line 282 · raw
@-n:Nat -> @-c:Nat -> @+cc:Nat -> @+nf:F32 -> @p:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, c> -> @labels:List<&2, Nat> -> F32
def ce_loss source · line 288 · raw
@+n:Nat -> @+c:Nat -> @logits:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, c> -> @labels:List<&2, Nat> -> F32
mean cross-entropy loss: -(1/n) Σ log p[i][label_i]
def ce_grad source · line 292 · raw
@+n:Nat -> @+c:Nat -> @logits:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, c> -> @labels:List<&2, Nat> -> 0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<n, c>
gradient of the loss with respect to the logits: (softmax - one-hot) / n
def sgd source · line 295 · raw
@-r:Nat -> @-c:Nat -> @+lr:F32 -> @w:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c> -> @dw:0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c> -> 0xf9837737d2c3f58ae1d0c5df42255584/main.Mat<r, c>
def sgd_vec source · line 298 · raw
@-n:Nat -> @+lr:F32 -> @b:0xf9837737d2c3f58ae1d0c5df42255584/main.Vec<n> -> @db:0xf9837737d2c3f58ae1d0c5df42255584/main.Vec<n> -> 0xf9837737d2c3f58ae1d0c5df42255584/main.Vec<n>