~/bend-docscommunity

main.bend checks

raw source on the hub · import bend-ml-tensor@0.1.2.0/main.bend as Main

bend-ml-tensor: tensors with the shape in the TYPE. A shape error is a type error.

import bend-ml-tensor@0.1.2.0/main.bend as T

Vec<n> is a vector of n F32 numbers; Mat<r, c> is an r x c matrix (a list of r rows with c numbers each). The dimensions are erased type parameters: they cost nothing at run time, but the checker verifies them in every operation.

2 imports
import Base
import bend-ml-nat-lemmas@0.1.0.0/main.bend as NL

Laws

law reshape_swap proved

Also proved in bend-mathlib as nat.mul_comm: import bend-mathlib@0.7.2.0/nat.bend as MNat, then MNat.mul_comm.

source · line 481 · raw

@+r:Nat -> @+c:Nat -> {Nat.mul(r, c) == Nat.mul(c, r) : Nat}

LAW (reshape_swap): r x c and c x r have the same number of elements (multiplication is commutative), so this reshape always exists.

law reshape_flat provedsource · line 491 · raw

@+r:Nat -> @+c:Nat -> {Nat.mul(r, c) == Nat.mul(1n, Nat.mul(r, c)) : Nat}

LAW (reshape_flat): r x c has the same number of elements as 1 x (r*c) (multiplying by 1 changes nothing), so flattening a matrix always exists.

Types

type Vec source · line 12 · raw

@-n:Nat -> Data

type Mat source · line 15 · raw

@-r:Nat -> @-c:Nat -> Data

Definitions

def zip_with.add source · line 22 · raw

@xs:List<&2, F32> -> @ys:List<&2, F32> -> List<&2, F32>

def zip_with.sub source · line 29 · raw

@xs:List<&2, F32> -> @ys:List<&2, F32> -> List<&2, F32>

def zip_with.mul source · line 36 · raw

@xs:List<&2, F32> -> @ys:List<&2, F32> -> List<&2, F32>

def scale.l source · line 43 · raw

@+k:F32 -> @xs:List<&2, F32> -> List<&2, F32>

def dot.go source · line 50 · raw

@xs:List<&2, F32> -> @ys:List<&2, F32> -> @acc:F32 -> F32

def sum.go source · line 57 · raw

@xs:List<&2, F32> -> @acc:F32 -> F32

def relu.l source · line 64 · raw

@xs:List<&2, F32> -> List<&2, F32>

def fill.l source · line 71 · raw

@n:Nat -> @+v:F32 -> List<&2, F32>

def Vec.zeros source · line 82 · raw

@n:Nat -> Vec<n>

def Vec.fill source · line 85 · raw

@n:Nat -> @+v:F32 -> Vec<n>

def Vec.add source · line 88 · raw

@-n:Nat -> @a:Vec<n> -> @b:Vec<n> -> Vec<n>

def Vec.sub source · line 93 · raw

@-n:Nat -> @a:Vec<n> -> @b:Vec<n> -> Vec<n>

def Vec.mul source · line 98 · raw

@-n:Nat -> @a:Vec<n> -> @b:Vec<n> -> Vec<n>

def Vec.scale source · line 103 · raw

@-n:Nat -> @+k:F32 -> @a:Vec<n> -> Vec<n>

def Vec.dot source · line 108 · raw

@-n:Nat -> @a:Vec<n> -> @b:Vec<n> -> F32

def Vec.sum source · line 113 · raw

@-n:Nat -> @a:Vec<n> -> F32

def Vec.relu source · line 118 · raw

@-n:Nat -> @a:Vec<n> -> Vec<n>

def heads source · line 127 · raw

@rows:List<&2, List<&2, F32>> -> List<&2, F32>

def tails source · line 134 · raw

@rows:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def transpose.go source · line 142 · raw

@c:Nat -> @+rows:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

transpose of a matrix with c columns (c is the one that shrinks, so it comes first)

def rows.zip_add source · line 149 · raw

@xs:List<&2, List<&2, F32>> -> @ys:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def rows.zip_sub source · line 156 · raw

@xs:List<&2, List<&2, F32>> -> @ys:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def rows.zip_mul source · line 163 · raw

@xs:List<&2, List<&2, F32>> -> @ys:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def rows.scale source · line 170 · raw

@+k:F32 -> @xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def rows.relu source · line 177 · raw

@xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def rows.col_sums source · line 185 · raw

@c:Nat -> @+rows:List<&2, List<&2, F32>> -> List<&2, F32>

column sums

def rows.add_row source · line 192 · raw

@xs:List<&2, List<&2, F32>> -> @+b:List<&2, F32> -> List<&2, List<&2, F32>>

def matmul.row source · line 199 · raw

@+ra:List<&2, F32> -> @bt:List<&2, List<&2, F32>> -> List<&2, F32>

def matmul.rows source · line 206 · raw

@a:List<&2, List<&2, F32>> -> @+bt:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def flatten.l source · line 214 · raw

@rows:List<&2, List<&2, F32>> -> List<&2, F32>

joins all the rows into a single list

def chunk.l source · line 222 · raw

@r:Nat -> @+c:Nat -> @+xs:List<&2, F32> -> List<&2, List<&2, F32>>

splits a list into r rows of c numbers

def Mat.zeros source · line 233 · raw

@+r:Nat -> @+c:Nat -> Mat<r, c>

def Mat.fill source · line 236 · raw

@+r:Nat -> @+c:Nat -> @+v:F32 -> Mat<r, c>

def Mat.add source · line 239 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> Mat<r, c>

def Mat.sub source · line 244 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> Mat<r, c>

def Mat.mul source · line 250 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> Mat<r, c>

element-wise product (Hadamard)

def Mat.scale source · line 255 · raw

@-r:Nat -> @-c:Nat -> @+k:F32 -> @a:Mat<r, c> -> Mat<r, c>

def Mat.relu source · line 260 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> Mat<r, c>

def Mat.transpose source · line 265 · raw

@-r:Nat -> @+c:Nat -> @a:Mat<r, c> -> Mat<c, r>

def Mat.matmul source · line 271 · raw

@-n:Nat -> @+k:Nat -> @+m:Nat -> @a:Mat<n, k> -> @b:Mat<k, m> -> Mat<n, m>

(n x k) · (k x m) = (n x m): the inner dimensions must be the same k

def Mat.matmul_t source · line 279 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> @a:Mat<n, k> -> @bt:Mat<m, k> -> Mat<n, m>

(n x k) · (m x k)ᵀ = (n x m): the second operand arrives already TRANSPOSED (stored as m rows of k numbers). It avoids transposing large weights on every call; the type still requires the same inner dimension k.

def Mat.add_row source · line 285 · raw

@-n:Nat -> @-m:Nat -> @a:Mat<n, m> -> @b:Vec<m> -> Mat<n, m>

adds a bias (a vector of m numbers) to each row of an n x m matrix

def Mat.col_sums source · line 291 · raw

@-n:Nat -> @+m:Nat -> @a:Mat<n, m> -> Vec<m>

sum of each column: Mat<n, m> -> Vec<m>

def Mat.matvec source · line 297 · raw

@-n:Nat -> @+k:Nat -> @a:Mat<n, k> -> @v:Vec<k> -> Vec<n>

matrix · vector: (n x k) · k = n

def Mat.reshape source · line 303 · raw

@-r1:Nat -> @-c1:Nat -> @r2:Nat -> @+c2:Nat -> @-p:{Nat.mul(r1, c1) == Nat.mul(r2, c2) : Nat} -> @a:Mat<r1, c1> -> Mat<r2, c2>

Reshape: compiles only with the PROOF that the number of elements does not change.

def max.go source · line 312 · raw

@xs:List<&2, F32> -> @cur:F32 -> F32

def max.l source · line 319 · raw

@xs:List<&2, F32> -> F32

def exp_shift source · line 326 · raw

@xs:List<&2, F32> -> @+m:F32 -> List<&2, F32>

def div.l source · line 333 · raw

@xs:List<&2, F32> -> @+s:F32 -> List<&2, F32>

def softmax.fin source · line 341 · raw

@+es:List<&2, F32> -> List<&2, F32>

stable softmax: subtracts the maximum before exponentiating

def softmax.l source · line 344 · raw

@+xs:List<&2, F32> -> List<&2, F32>

def rows.softmax source · line 347 · raw

@xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def gelu.f source · line 355 · raw

@+x:F32 -> F32

GELU with the tanh approximation (the same as GPT-2)

def gelu.l source · line 358 · raw

@xs:List<&2, F32> -> List<&2, F32>

def rows.gelu source · line 365 · raw

@xs:List<&2, List<&2, F32>> -> List<&2, List<&2, F32>>

def sub_const source · line 372 · raw

@xs:List<&2, F32> -> @+k:F32 -> List<&2, F32>

def sq_sum source · line 379 · raw

@xs:List<&2, F32> -> @acc:F32 -> F32

def layernorm.fin source · line 387 · raw

@+cs:List<&2, F32> -> @+d:F32 -> @+eps:F32 -> @+gamma:List<&2, F32> -> @+beta:List<&2, F32> -> List<&2, F32>

(x - mean) / sqrt(variance + eps), then * gamma + beta

def layernorm.l source · line 390 · raw

@+xs:List<&2, F32> -> @+d:F32 -> @+eps:F32 -> @+gamma:List<&2, F32> -> @+beta:List<&2, F32> -> List<&2, F32>

def rows.layernorm source · line 393 · raw

@xs:List<&2, List<&2, F32>> -> @+d:F32 -> @+eps:F32 -> @+gamma:List<&2, F32> -> @+beta:List<&2, F32> -> List<&2, List<&2, F32>>

def Mat.softmax source · line 400 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> Mat<r, c>

def Mat.gelu source · line 405 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> Mat<r, c>

def Mat.layernorm source · line 411 · raw

@-n:Nat -> @d:Nat -> @eps:F32 -> @g:Vec<d> -> @bt:Vec<d> -> @a:Mat<n, d> -> Mat<n, d>

normalizes each row (d columns) and applies gamma and beta, vectors of d numbers

def Vec.softmax source · line 416 · raw

@-n:Nat -> @a:Vec<n> -> Vec<n>

def len.eq source · line 425 · raw

@xs:List<&2, F32> -> @n:Nat -> Bool

def rows.ok source · line 434 · raw

@rs:List<&2, List<&2, F32>> -> @+c:Nat -> @r:Nat -> Bool

def Vec.of2 source · line 443 · raw

@-n:Nat -> @ok:Bool -> @xs:List<&2, F32> -> Maybe<&2, Vec<n>>

def Vec.of source · line 451 · raw

@n:Nat -> @+xs:List<&2, F32> -> Maybe<&2, Vec<n>>

Vec<n> from a list, only if it has exactly n numbers

def Mat.of2 source · line 454 · raw

@-r:Nat -> @-c:Nat -> @ok:Bool -> @rs:List<&2, List<&2, F32>> -> Maybe<&2, Mat<r, c>>

def Mat.of source · line 462 · raw

@+r:Nat -> @+c:Nat -> @+rs:List<&2, List<&2, F32>> -> Maybe<&2, Mat<r, c>>

Mat<r, c> from rows, only if they are exactly r rows of c numbers

def Vec.to_list source · line 465 · raw

@-n:Nat -> @a:Vec<n> -> List<&2, F32>

def Mat.to_rows source · line 470 · raw

@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> List<&2, List<&2, F32>>

def Mat.flatten source · line 500 · raw

@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> Mat<1n, Nat.mul(r, c)>

Mat<r, c> -> Mat<1, r*c>: the proof comes from the law above

def Mat.reshape_swap source · line 504 · raw

@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> Mat<c, r>

Mat<r, c> -> Mat<c, r> over the same data (it is NOT the transpose!): it only reinterprets