~/bend-docscommunity

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).

type G source · line 126 · raw

Data

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>