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.
@+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
Vec@-n:Nat -> @xs:List<&2, F32> -> Vec<n>
type Mat source · line 15 · raw
@-r:Nat -> @-c:Nat -> Data
Mat@-r:Nat -> @-c:Nat -> @rows:List<&2, List<&2, F32>> -> Mat<r, c>
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