main.bend checks
raw source on the hub · import bend-ml-tensor-array@0.1.2.0/main.bend as Main
bend-ml-tensor-array: tensors over a flat Array<F32>, with the shape in the type.
import bend-ml-tensor-array@0.1.2.0/main.bend as TA
Mat<r, c> stores r*c numbers in an Array<F32>, row by row (index i*c + j). The dimensions are erased type parameters: the same guarantee as bend-ml-tensor (a product with wrong dimensions does not compile), but ~50x faster than lists: Array.get/set by index, and products in parallel blocks of rows.
An Array is affine (a single owner): every operation that reads an operand returns it together with the result (as Array.get does), in records MM, AR, ... Use Mat.clone when you need two copies.
What is NOT proved: the invariant "the Array has capacity >= r*c" holds because the constructors (Mat.zeros, Mat.of_list) allocate the right size; and F32 is validated by tests against PyTorch (reference/test_tensor_array.py), not by proof.
1 import
import Base
Types
type R source · line 24 · raw
Type
R@a:Array<F32> -> @b:Array<F32> -> @x:F32 -> R
type G source · line 27 · raw
Type
G@a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
type S3 source · line 66 · raw
Type
S3@a:Array<F32> -> @b:Array<F32> -> @out:List<&2, F32> -> S3
type P2 source · line 176 · raw
Type
P2@a:Array<F32> -> @b:Array<F32> -> P2
type AL source · line 179 · raw
Type
AL@c:Array<F32> -> @l:List<&2, F32> -> AL
type L2 source · line 318 · raw
Type
L2@loss:F32 -> @c:Array<F32> -> L2
type Mat source · line 331 · raw
@-r:Nat -> @-c:Nat -> Type
Mat@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>
type CR source · line 345 · raw
Type
CR@c:Array<F32> -> @n:Nat -> CR
type ML source · line 449 · raw
@-r:Nat -> @-c:Nat -> Type
ML@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> @l:List<&2, F32> -> ML<r, c>
type MM2 source · line 464 · raw
@-r:Nat -> @-c:Nat -> Type
MM2@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> MM2<r, c>
type MMul source · line 481 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> Type
MMul@-n:Nat -> @-k:Nat -> @-m:Nat -> @a:Mat<n, k> -> @b:Mat<k, m> -> @c:Mat<n, m> -> MMul<n, k, m>
type MMulNT source · line 486 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> Type
MMulNT@-n:Nat -> @-k:Nat -> @-m:Nat -> @a:Mat<n, k> -> @b:Mat<m, k> -> @c:Mat<n, m> -> MMulNT<n, k, m>
type MMulTN source · line 491 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> Type
MMulTN@-n:Nat -> @-k:Nat -> @-m:Nat -> @a:Mat<k, n> -> @b:Mat<k, m> -> @c:Mat<n, m> -> MMulTN<n, k, m>
type MBias source · line 560 · raw
@-n:Nat -> @-m:Nat -> Type
MBias@-n:Nat -> @-m:Nat -> @c:Mat<n, m> -> @b:Mat<1n, m> -> MBias<n, m>
type MPair source · line 575 · raw
@-r:Nat -> @-c:Nat -> Type
MPair@-r:Nat -> @-c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> MPair<r, c>
type MColSum source · line 614 · raw
@-n:Nat -> @-m:Nat -> Type
MColSum@-n:Nat -> @-m:Nat -> @a:Mat<n, m> -> @s:Mat<1n, m> -> MColSum<n, m>
type MCE source · line 630 · raw
@-n:Nat -> @-c:Nat -> Type
MCE@-n:Nat -> @-c:Nat -> @loss:F32 -> @grad:Mat<n, c> -> MCE<n, c>
type MHits source · line 645 · raw
@-n:Nat -> @-c:Nat -> Type
MHits@-n:Nat -> @-c:Nat -> @z:Mat<n, c> -> @hits:Nat -> MHits<n, c>
type MRow source · line 660 · raw
@-r:Nat -> @-c:Nat -> Type
MRow@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> @row:List<&2, F32> -> MRow<r, c>
Definitions
def dotg source · line 32 · raw
@k:Nat -> @+ia:U32 -> @+ib:U32 -> @+sa:U32 -> @+sb:U32 -> @acc:F32 -> @ra:Pair(Array<F32>, F32) -> @rb:Pair(Array<F32>, F32) -> R
def gcols source · line 41 · raw
@mleft:Nat -> @+kn:Nat -> @+ia0:U32 -> @+ib0:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @+cidx:U32 -> @c:Array<F32> -> @r:R -> G
def grows source · line 48 · raw
@nleft:Nat -> @+mm1:Nat -> @+kn:Nat -> @+i:U32 -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @+mu:U32 -> @g:G -> G
def gemm source · line 57 · raw
@+n:Nat -> @+m:Nat -> @+k:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def lcols source · line 69 · raw
@mleft:Nat -> @+kn:Nat -> @+ia0:U32 -> @+ib0:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @out:List<&2, F32> -> @r:R -> S3
def lrows source · line 76 · raw
@nleft:Nat -> @+mm1:Nat -> @+kn:Nat -> @+i:U32 -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @st:S3 -> S3
def lblock_fin source · line 83 · raw
@st:S3 -> List<&2, F32>
def lblock source · line 88 · raw
@cnt:Nat -> @+mm1:Nat -> @+kn:Nat -> @+i:U32 -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> List<&2, F32>
def lpar source · line 91 · raw
@d:Nat -> @+cnt:Nat -> @+mm1:Nat -> @+kn:Nat -> @+i:U32 -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @ra:Pair(Array<F32>, Array<F32>) -> @rb:Pair(Array<F32>, Array<F32>) -> List<&2, F32>
def write_l source · line 99 · raw
@xs:List<&2, F32> -> @+idx:U32 -> @c:Array<F32> -> Array<F32>
def gp3 source · line 108 · raw
@a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> @xs:List<&2, F32> -> G
def gp2 source · line 111 · raw
@d:Nat -> @+n:Nat -> @+mm1:Nat -> @+kn:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @c:Array<F32> -> @pa:Pair(Array<F32>, Array<F32>) -> @pb:Pair(Array<F32>, Array<F32>) -> G
def gemm_par source · line 118 · raw
@d:Nat -> @+n:Nat -> @+m:Nat -> @+k:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def cblock_fin source · line 127 · raw
@st:S3 -> List<&2, F32>
def cblock_go source · line 134 · raw
@cnt1:Nat -> @+kn:Nat -> @+ia0:U32 -> @+j0:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> List<&2, F32>
def cblock source · line 137 · raw
@cnt:Nat -> @+kn:Nat -> @+ia0:U32 -> @+j0:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> List<&2, F32>
def cpar source · line 144 · raw
@d:Nat -> @+cnt:Nat -> @+kn:Nat -> @+ia0:U32 -> @+j0:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @ra:Pair(Array<F32>, Array<F32>) -> @rb:Pair(Array<F32>, Array<F32>) -> List<&2, F32>
def cp2 source · line 152 · raw
@d:Nat -> @+m:Nat -> @+kn:Nat -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @c:Array<F32> -> @pa:Pair(Array<F32>, Array<F32>) -> @pb:Pair(Array<F32>, Array<F32>) -> G
def gemm_par_cols source · line 159 · raw
@d:Nat -> @+m:Nat -> @+k:Nat -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def depth source · line 166 · raw
@n:Nat -> Nat
def read_l source · line 184 · raw
@mleft:Nat -> @+idx:U32 -> @acc:List<&2, F32> -> @rc:Pair(Array<F32>, F32) -> AL
def addl source · line 193 · raw
@xs:List<&2, F32> -> @+idx:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>
def bias_rows source · line 202 · raw
@nleft:Nat -> @+i:U32 -> @+mu:U32 -> @+bl:List<&2, F32> -> @c:Array<F32> -> Array<F32>
def relu_l source · line 211 · raw
@nleft:Nat -> @+t:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>
def pos2 source · line 218 · raw
@b:Bool -> F32
def pos source · line 225 · raw
@+y:F32 -> F32
def mask_l source · line 230 · raw
@nleft:Nat -> @+t:U32 -> @rd:Pair(Array<F32>, F32) -> @rh:Pair(Array<F32>, F32) -> P2
def maxl2 source · line 239 · raw
@xs:List<&2, F32> -> @cur:F32 -> F32
def maxl source · line 246 · raw
@xs:List<&2, F32> -> F32
def expl source · line 253 · raw
@xs:List<&2, F32> -> @+m:F32 -> List<&2, F32>
def suml source · line 260 · raw
@xs:List<&2, F32> -> @acc:F32 -> F32
def divl source · line 267 · raw
@xs:List<&2, F32> -> @+s:F32 -> List<&2, F32>
def softmax2 source · line 274 · raw
@+es:List<&2, F32> -> List<&2, F32>
def softmax source · line 277 · raw
@+xs:List<&2, F32> -> List<&2, F32>
def sgd_l source · line 282 · raw
@nleft:Nat -> @+t:U32 -> @+lr:F32 -> @rw:Pair(Array<F32>, F32) -> @rd:Pair(Array<F32>, F32) -> P2
def pick source · line 293 · raw
@ps:List<&2, F32> -> @n:Nat -> F32
def hit source · line 304 · raw
@b:Bool -> F32
def grad_l source · line 311 · raw
@ps:List<&2, F32> -> @+label:Nat -> @+j:Nat -> @+nf:F32 -> List<&2, F32>
def ce_row2 source · line 321 · raw
@+ps:List<&2, F32> -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> @c:Array<F32> -> L2
def ce_row source · line 324 · raw
@al:AL -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> L2
def ce_rows source · line 336 · raw
@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+nf:F32 -> @+i:U32 -> @r:L2 -> L2
def gt_head source · line 348 · raw
@xs:List<&2, F32> -> @+v:F32 -> Bool
def argmax2 source · line 355 · raw
@xs:List<&2, F32> -> @g:Bool -> @+i:Nat -> @+bi:Nat -> @+bv:F32 -> Nat
def argmax source · line 364 · raw
@xs:List<&2, F32> -> Nat
def one source · line 371 · raw
@b:Bool -> Nat
def count_rows3 source · line 378 · raw
@+y:Nat -> @n:Nat -> @al:AL -> CR
def count_rows source · line 383 · raw
@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+i:U32 -> @r:CR -> CR
def cap_depth2 source · line 395 · raw
@small:Bool -> @n:Nat -> Nat
def cap_depth source · line 402 · raw
@+n:Nat -> Nat
def Mat.zeros source · line 407 · raw
@+r:Nat -> @+c:Nat -> Mat<r, c>
def Mat.fill source · line 410 · raw
@+r:Nat -> @+c:Nat -> @+v:F32 -> Mat<r, c>
def fill_list source · line 413 · raw
@xs:List<&2, F32> -> @+i:U32 -> @a:Array<F32> -> Array<F32>
def len_is source · line 420 · raw
@xs:List<&2, F32> -> @n:Nat -> Bool
def Mat.of_list2 source · line 429 · raw
@-r:Nat -> @-c:Nat -> @+rr:Nat -> @+cc:Nat -> @ok:Bool -> @xs:List<&2, F32> -> Maybe<&1, Mat<r, c>>
def Mat.of_list source · line 438 · raw
@+r:Nat -> @+c:Nat -> @+xs:List<&2, F32> -> Maybe<&1, Mat<r, c>>
def Mat.from_list source · line 446 · raw
@+r:Nat -> @+c:Nat -> @xs:List<&2, F32> -> Mat<r, c>
def Mat.to_list2 source · line 452 · raw
@-r:Nat -> @-c:Nat -> @al:AL -> ML<r, c>
def Mat.to_list source · line 459 · raw
@+r:Nat -> @+c:Nat -> @m:Mat<r, c> -> ML<r, c>
def Mat.clone2 source · line 467 · raw
@-r:Nat -> @-c:Nat -> @p:Pair(Array<F32>, Array<F32>) -> MM2<r, c>
def Mat.clone source · line 476 · raw
@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> MM2<r, c>
def gemm_any3 source · line 494 · raw
@cols:Bool -> @d:Nat -> @+n:Nat -> @+m:Nat -> @+k:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def gemm_any2 source · line 501 · raw
@d:Nat -> @+n:Nat -> @+m:Nat -> @+k:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def gemm_any source · line 508 · raw
@d:Nat -> @+n:Nat -> @+m:Nat -> @+k:Nat -> @+arow:U32 -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G
def mmul_fin source · line 511 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMul<n, k, m>
def mmul_nt_fin source · line 516 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulNT<n, k, m>
def mmul_tn_fin source · line 521 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulTN<n, k, m>
def mm_a source · line 526 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @d:Nat -> @a:Array<F32> -> @b:Array<F32> -> MMul<n, k, m>
def Mat.matmul source · line 531 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @par:Nat -> @a:Mat<n, k> -> @b:Mat<k, m> -> MMul<n, k, m>
def mm_nt source · line 536 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @d:Nat -> @a:Array<F32> -> @b:Array<F32> -> MMulNT<n, k, m>
def Mat.matmul_nt source · line 541 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @par:Nat -> @a:Mat<n, k> -> @b:Mat<m, k> -> MMulNT<n, k, m>
def mm_tn source · line 546 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @d:Nat -> @a:Array<F32> -> @b:Array<F32> -> MMulTN<n, k, m>
def Mat.matmul_tn source · line 551 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @par:Nat -> @a:Mat<k, n> -> @b:Mat<k, m> -> MMulTN<n, k, m>
def add_row2 source · line 563 · raw
@-n:Nat -> @-m:Nat -> @+nn:Nat -> @+mm:Nat -> @cd:Array<F32> -> @al:AL -> MBias<n, m>
def Mat.add_row source · line 568 · raw
@+n:Nat -> @+m:Nat -> @c:Mat<n, m> -> @b:Mat<1n, m> -> MBias<n, m>
def mpair_fin source · line 578 · raw
@-r:Nat -> @-c:Nat -> @p:P2 -> MPair<r, c>
def relu_fin source · line 583 · raw
@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>
def Mat.relu source · line 586 · raw
@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> Mat<r, c>
def Mat.relu_bwd source · line 593 · raw
@+r:Nat -> @+c:Nat -> @dy:Mat<r, c> -> @h:Mat<r, c> -> MPair<r, c>
def Mat.sgd source · line 600 · raw
@+r:Nat -> @+c:Nat -> @+lr:F32 -> @w:Mat<r, c> -> @dw:Mat<r, c> -> MPair<r, c>
def Mat.add source · line 607 · raw
@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> MPair<r, c>
def colsum_fin source · line 617 · raw
@-n:Nat -> @-m:Nat -> @g:G -> MColSum<n, m>
def Mat.col_sums source · line 622 · raw
@+n:Nat -> @+m:Nat -> @a:Mat<n, m> -> MColSum<n, m>
def ce_fin source · line 633 · raw
@-n:Nat -> @-c:Nat -> @+nf:F32 -> @r:L2 -> MCE<n, c>
def Mat.softmax_ce source · line 638 · raw
@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MCE<n, c>
def hits_fin source · line 648 · raw
@-n:Nat -> @-c:Nat -> @r:CR -> MHits<n, c>
def Mat.count_correct source · line 653 · raw
@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MHits<n, c>
def row_fin source · line 663 · raw
@-r:Nat -> @-c:Nat -> @al:AL -> MRow<r, c>
def Mat.read_row source · line 668 · raw
@+r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> MRow<r, c>
def Mat.write_row source · line 673 · raw
@-r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> @xs:List<&2, F32> -> Mat<r, c>