~/bend-docscommunity

main.bend checks

raw source on the hub · import bend-ml-tensor-array@0.1.3.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

Laws

law cap_ok provedsource · line 551 · raw

@+n:Nat -> {True{} == leb(n, pow2(cap_depth(n))) : Bool}

LAW (cap_ok): an Array with 2^cap_depth(n) slots always has room for n elements.

Types

type R source · line 24 · raw

Type

type G source · line 27 · raw

Type

type S3 source · line 64 · raw

Type

type P2 source · line 171 · raw

Type

type AL source · line 174 · raw

Type

type L2 source · line 313 · raw

Type

type Mat source · line 326 · raw

@-r:Nat -> @-c:Nat -> Type

type CR source · line 338 · raw

Type

type ML source · line 609 · raw

@-r:Nat -> @-c:Nat -> Type

type MM2 source · line 624 · raw

@-r:Nat -> @-c:Nat -> Type

type MMul source · line 641 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> Type

type MMulNT source · line 646 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> Type

type MMulTN source · line 651 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> Type

type MBias source · line 720 · raw

@-n:Nat -> @-m:Nat -> Type

type MPair source · line 735 · raw

@-r:Nat -> @-c:Nat -> Type

type MColSum source · line 774 · raw

@-n:Nat -> @-m:Nat -> Type

type MCE source · line 790 · raw

@-n:Nat -> @-c:Nat -> Type

type MHits source · line 805 · raw

@-n:Nat -> @-c:Nat -> Type

type MRow source · line 820 · raw

@-r:Nat -> @-c:Nat -> Type

Definitions

def dotg source · line 30 · 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 39 · 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 46 · 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 55 · 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 67 · 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 74 · 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 81 · raw

@st:S3 -> List<&2, F32>

def lblock source · line 86 · 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 89 · 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 97 · raw

@xs:List<&2, F32> -> @+idx:U32 -> @c:Array<F32> -> Array<F32>

def gp3 source · line 106 · raw

@a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> @xs:List<&2, F32> -> G

def gp2 source · line 109 · 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 116 · 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 125 · raw

@st:S3 -> List<&2, F32>

def cblock_go source · line 132 · 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 135 · 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 142 · 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 150 · 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 157 · raw

@d:Nat -> @+m:Nat -> @+k:Nat -> @+sa:U32 -> @+sb:U32 -> @+bcol:U32 -> @a:Array<F32> -> @b:Array<F32> -> @c:Array<F32> -> G

def read_l source · line 179 · raw

@mleft:Nat -> @+idx:U32 -> @acc:List<&2, F32> -> @rc:Pair(Array<F32>, F32) -> AL

def addl source · line 188 · raw

@xs:List<&2, F32> -> @+idx:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>

def bias_rows source · line 197 · raw

@nleft:Nat -> @+i:U32 -> @+mu:U32 -> @+bl:List<&2, F32> -> @c:Array<F32> -> Array<F32>

def relu_l source · line 206 · raw

@nleft:Nat -> @+t:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>

def pos2 source · line 213 · raw

@b:Bool -> F32

def pos source · line 220 · raw

@+y:F32 -> F32

def mask_l source · line 225 · raw

@nleft:Nat -> @+t:U32 -> @rd:Pair(Array<F32>, F32) -> @rh:Pair(Array<F32>, F32) -> P2

def maxl2 source · line 234 · raw

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

def maxl source · line 241 · raw

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

def expl source · line 248 · raw

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

def suml source · line 255 · raw

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

def divl source · line 262 · raw

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

def softmax2 source · line 269 · raw

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

def softmax source · line 272 · raw

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

def sgd_l source · line 277 · raw

@nleft:Nat -> @+t:U32 -> @+lr:F32 -> @rw:Pair(Array<F32>, F32) -> @rd:Pair(Array<F32>, F32) -> P2

def pick source · line 288 · raw

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

def hit source · line 299 · raw

@b:Bool -> F32

def grad_l source · line 306 · raw

@ps:List<&2, F32> -> @+label:Nat -> @+j:Nat -> @+nf:F32 -> List<&2, F32>

def ce_row2 source · line 316 · raw

@+ps:List<&2, F32> -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> @c:Array<F32> -> L2

def ce_row source · line 319 · raw

@al:AL -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> L2

def ce_rows source · line 329 · raw

@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+nf:F32 -> @+i:U32 -> @r:L2 -> L2

def gt_head source · line 341 · raw

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

def argmax2 source · line 348 · raw

@xs:List<&2, F32> -> @g:Bool -> @+i:Nat -> @+bi:Nat -> @+bv:F32 -> Nat

def argmax source · line 357 · raw

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

def one source · line 364 · raw

@b:Bool -> Nat

def count_rows3 source · line 371 · raw

@+y:Nat -> @n:Nat -> @al:AL -> CR

def count_rows source · line 376 · raw

@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+i:U32 -> @r:CR -> CR

def BD source · line 397 · raw

@b:Bool -> @t:Type -> @f:Type -> Type

a type chosen by a Bool: rewriting through it refutes True == False

def leb source · line 405 · raw

@a:Nat -> @b:Nat -> Bool

a <= b, by direct recursion

def pow2 source · line 414 · raw

@d:Nat -> Nat

def cap_pick source · line 422 · raw

@ok:Bool -> @d:Nat -> @+p:Nat -> Pair(Nat, Nat)

smallest power of two >= n, by structural recursion on n: returns (depth, 2^depth)

def cap_step source · line 429 · raw

@r:Pair(Nat, Nat) -> @+n:Nat -> Pair(Nat, Nat)

def cap_pair source · line 434 · raw

@n:Nat -> Pair(Nat, Nat)

def fst_nat source · line 441 · raw

@r:Pair(Nat, Nat) -> Nat

def snd_nat source · line 446 · raw

@r:Pair(Nat, Nat) -> Nat

def cap_depth source · line 452 · raw

@n:Nat -> Nat

the depth d such that an Array with 2^d slots holds n elements

def leb_succ source · line 458 · raw

@a:Nat -> @b:Nat -> @e:{True{} == leb(a, b) : Bool} -> {True{} == leb(a, 1n+b) : Bool}

a <= b -> a <= 1 + b

def leb_up_l source · line 469 · raw

@+a:Nat -> @+b:Nat -> @c:Nat -> @+e:{True{} == leb(a, b) : Bool} -> {True{} == leb(a, Nat.add(c, b)) : Bool}

a <= b -> a <= c + b

def leb_up source · line 477 · raw

@a:Nat -> @b:Nat -> @c:Nat -> @e:{True{} == leb(a, b) : Bool} -> {True{} == leb(a, Nat.add(b, c)) : Bool}

a <= b -> a <= b + c

def leb_add source · line 488 · raw

@a:Nat -> @b:Nat -> @c:Nat -> @d:Nat -> @e1:{True{} == leb(a, b) : Bool} -> @e2:{True{} == leb(c, d) : Bool} -> {True{} == leb(Nat.add(a, c), Nat.add(b, d)) : Bool}

a <= b and c <= d -> a + c <= b + d

def pow2_pos source · line 499 · raw

@d:Nat -> {True{} == leb(1n, pow2(d)) : Bool}

1 <= 2^d

def pos_p source · line 507 · raw

@d:Nat -> @p:Nat -> @i1:{p == pow2(d) : Nat} -> {True{} == leb(1n, p) : Bool}

1 <= p when p = 2^d

def step_eq source · line 512 · raw

@d:Nat -> @p:Nat -> @i1:{p == pow2(d) : Nat} -> {Nat.add(p, p) == pow2(1n+d) : Nat}

p + p == 2^(1+d) when p == 2^d

def step_ok source · line 517 · raw

@ok:Bool -> @+d:Nat -> @+p:Nat -> @+m:Nat -> @+e:{ok == leb(1n+m, p) : Bool} -> @+i1:{p == pow2(d) : Nat} -> @+i2:{True{} == leb(m, p) : Bool} -> Pair({snd_nat(cap_pick(ok, d, p)) == pow2(fst_nat(cap_pick(ok, d, p))) : Nat}, {True{} == leb(1n+m, snd_nat(cap_pick(ok, d, p))) : Bool})

one step: given the invariant for m, the invariant for 1+m holds whichever way the comparison goes

def step_pair_ok source · line 526 · raw

@+m:Nat -> @r:Pair(Nat, Nat) -> @i1:{snd_nat(r) == pow2(fst_nat(r)) : Nat} -> @i2:{True{} == leb(m, snd_nat(r)) : Bool} -> Pair({snd_nat(cap_step(r, 1n+m)) == pow2(fst_nat(cap_step(r, 1n+m))) : Nat}, {True{} == leb(1n+m, snd_nat(cap_step(r, 1n+m))) : Bool})

the invariant for cap_step on a pair, opened

def cap_succ source · line 531 · raw

@+m:Nat -> @ih:Pair({snd_nat(cap_pair(m)) == pow2(fst_nat(cap_pair(m))) : Nat}, {True{} == leb(m, snd_nat(cap_pair(m))) : Bool}) -> Pair({snd_nat(cap_pair(1n+m)) == pow2(fst_nat(cap_pair(1n+m))) : Nat}, {True{} == leb(1n+m, snd_nat(cap_pair(1n+m))) : Bool})

def cap_inv source · line 537 · raw

@n:Nat -> Pair({snd_nat(cap_pair(n)) == pow2(fst_nat(cap_pair(n))) : Nat}, {True{} == leb(n, snd_nat(cap_pair(n))) : Bool})

the invariant of cap_pair: the second component is 2^(first) and is >= n

def cap_final source · line 544 · raw

@n:Nat -> @pr:Pair({snd_nat(cap_pair(n)) == pow2(fst_nat(cap_pair(n))) : Nat}, {True{} == leb(n, snd_nat(cap_pair(n))) : Bool}) -> {True{} == leb(n, pow2(cap_depth(n))) : Bool}

def Mat.zeros source · line 560 · raw

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

def Mat.fill source · line 563 · raw

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

def fill_list source · line 566 · raw

@xs:List<&2, F32> -> @+i:U32 -> @a:Array<F32> -> Array<F32>

def len_is source · line 573 · raw

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

def Mat.of_list2 source · line 582 · 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 591 · raw

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

def Mat.from_list source · line 599 · raw

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

def Mat.fill_at source · line 604 · raw

@-r:Nat -> @-c:Nat -> @+i:U32 -> @xs:List<&2, F32> -> @m:Mat<r, c> -> Mat<r, c>

writes the numbers of xs into m starting at flat index i (row-major); numbers past the capacity are dropped. For loading a large matrix block by block without building one big list.

def Mat.to_list2 source · line 612 · raw

@-r:Nat -> @-c:Nat -> @al:AL -> ML<r, c>

def Mat.to_list source · line 619 · raw

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

def Mat.clone2 source · line 627 · raw

@-r:Nat -> @-c:Nat -> @p:Pair(Array<F32>, Array<F32>) -> MM2<r, c>

def Mat.clone source · line 636 · raw

@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> MM2<r, c>

def gemm_any3 source · line 654 · 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 661 · 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 668 · 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 671 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMul<n, k, m>

def mmul_nt_fin source · line 676 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulNT<n, k, m>

def mmul_tn_fin source · line 681 · raw

@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulTN<n, k, m>

def mm_a source · line 686 · raw

@+n:Nat -> @+k:Nat -> @+m:Nat -> @d:Nat -> @a:Array<F32> -> @b:Array<F32> -> MMul<n, k, m>

def Mat.matmul source · line 691 · 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 696 · 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 701 · 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 706 · 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 711 · 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 723 · raw

@-n:Nat -> @-m:Nat -> @+nn:Nat -> @+mm:Nat -> @cd:Array<F32> -> @al:AL -> MBias<n, m>

def Mat.add_row source · line 728 · raw

@+n:Nat -> @+m:Nat -> @c:Mat<n, m> -> @b:Mat<1n, m> -> MBias<n, m>

def mpair_fin source · line 738 · raw

@-r:Nat -> @-c:Nat -> @p:P2 -> MPair<r, c>

def relu_fin source · line 743 · raw

@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>

def Mat.relu source · line 746 · raw

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

def Mat.relu_bwd source · line 753 · raw

@+r:Nat -> @+c:Nat -> @dy:Mat<r, c> -> @h:Mat<r, c> -> MPair<r, c>

def Mat.sgd source · line 760 · raw

@+r:Nat -> @+c:Nat -> @+lr:F32 -> @w:Mat<r, c> -> @dw:Mat<r, c> -> MPair<r, c>

def Mat.add source · line 767 · raw

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

def colsum_fin source · line 777 · raw

@-n:Nat -> @-m:Nat -> @g:G -> MColSum<n, m>

def Mat.col_sums source · line 782 · raw

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

def ce_fin source · line 793 · raw

@-n:Nat -> @-c:Nat -> @+nf:F32 -> @r:L2 -> MCE<n, c>

def Mat.softmax_ce source · line 798 · raw

@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MCE<n, c>

def hits_fin source · line 808 · raw

@-n:Nat -> @-c:Nat -> @r:CR -> MHits<n, c>

def Mat.count_correct source · line 813 · raw

@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MHits<n, c>

def row_fin source · line 823 · raw

@-r:Nat -> @-c:Nat -> @al:AL -> MRow<r, c>

def Mat.read_row source · line 828 · raw

@+r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> MRow<r, c>

def Mat.write_row source · line 833 · raw

@-r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> @xs:List<&2, F32> -> Mat<r, c>

def labels_len_is source · line 841 · raw

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

does the list of labels have exactly n entries (one per row)?

def ce_checked source · line 850 · raw

@ok:Bool -> @+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> Maybe<&1, MCE<n, c>>

def Mat.softmax_ce_checked source · line 858 · raw

@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @+labels:List<&2, Nat> -> Maybe<&1, MCE<n, c>>

like Mat.softmax_ce, but returns None unless labels has exactly n entries

def hits_checked source · line 861 · raw

@ok:Bool -> @+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> Maybe<&1, MHits<n, c>>

def Mat.count_correct_checked source · line 869 · raw

@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @+labels:List<&2, Nat> -> Maybe<&1, MHits<n, c>>

like Mat.count_correct, but returns None unless labels has exactly n entries