main.bend checks
raw source on the hub · import bend-ml-tensor-array@0.1.6.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.6.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.
law half_cover provedsource · line 962 · raw
@+n:Nat -> {n == Nat.add(half(n), Nat.sub(n, half(n))) : Nat}LAW (half_cover): the two bands of a node, half(r) rows and r - half(r) rows, add up to exactly r rows: no row is lost and none is counted twice.
law leaf_len provedsource · line 1216 · raw
@-r:Nat -> @-c:Nat -> @w:Array<F32> -> @+mleft:Nat -> @+idx:U32 -> @+acc:List<&2, F32> -> @rc:Pair(Array<F32>, F32) -> {bl_len(r, c, bv_g2(r, c, w, read_l(mleft, idx, acc, rc))) == Nat.add(1n+mleft, List.length(&2, F32, acc)) : Nat}a leaf's list: read_l reads mleft + 1 numbers onto acc and reverses them. By induction on mleft: with 0 left it reads one number (rev_len counts the reversed list); with 1 + p left it reads one and recurses, and len_step moves the 1 out of the sum.
law band_len provedsource · line 1242 · raw
@-r:Nat -> @+c:Nat -> @+rows:Nat -> @x:Array<F32> -> @w:Array<F32> -> {bl_len(r, c, bv_gemm(r, c, rows, rows, x, w)) == rows : Nat}LAW (band_len): the product of one band with rows rows gives exactly rows numbers.
0 rows give the empty list; 1 + p rows go through the gemm, and g_len counts what is read back.
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 64 · raw
Type
S3@a:Array<F32> -> @b:Array<F32> -> @out:List<&2, F32> -> S3
type P2 source · line 171 · raw
Type
P2@a:Array<F32> -> @b:Array<F32> -> P2
type AL source · line 174 · raw
Type
AL@c:Array<F32> -> @l:List<&2, F32> -> AL
type L2 source · line 313 · raw
Type
L2@loss:F32 -> @c:Array<F32> -> L2
type Mat source · line 326 · raw
@-r:Nat -> @-c:Nat -> Type
Mat@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>
type CR source · line 338 · raw
Type
CR@c:Array<F32> -> @n:Nat -> CR
type ML source · line 609 · 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 624 · 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 641 · 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 646 · 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 651 · 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 720 · 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 735 · 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 774 · 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 790 · raw
@-n:Nat -> @-c:Nat -> Type
MCE@-n:Nat -> @-c:Nat -> @loss:F32 -> @grad:Mat<n, c> -> MCE<n, c>
type MHits source · line 805 · raw
@-n:Nat -> @-c:Nat -> Type
MHits@-n:Nat -> @-c:Nat -> @z:Mat<n, c> -> @hits:Nat -> MHits<n, c>
type MRow source · line 820 · raw
@-r:Nat -> @-c:Nat -> Type
MRow@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> @row:List<&2, F32> -> MRow<r, c>
type Bands source · line 969 · raw
@-r:Nat -> @-c:Nat -> Type
BLeaf@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> Bands<r, c>
BNode@-r:Nat -> @-c:Nat -> @x:Bands<half(r), c> -> @y:Bands<Nat.sub(r, half(r)), c> -> Bands<r, c>
type BV source · line 1041 · raw
@-r:Nat -> @-c:Nat -> Type
B (intact) and y = B · x, typed like Mat.matmul_nt: x is 1 x c, y is 1 x r
BV@-r:Nat -> @-c:Nat -> @b:Bands<r, c> -> @y:Mat<1n, r> -> BV<r, c>
type BL source · line 1045 · raw
@-r:Nat -> @-c:Nat -> Type
the bands and y as a list in row order (the result of Bands.matvec_l)
BL@-r:Nat -> @-c:Nat -> @b:Bands<r, c> -> @y:List<&2, F32> -> BL<r, c>
type BRow source · line 1110 · raw
@-r:Nat -> @-c:Nat -> Type
BRow@-r:Nat -> @-c:Nat -> @b:Bands<r, c> -> @row:List<&2, F32> -> BRow<r, c>
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
def half.go source · line 888 · raw
@n:Nat -> @acc:Nat -> Nat
n / 2 rounded down. A tail loop with an accumulator: a non-tail recursion here would put a continuation on every node of the band tree and keep the bands from running in parallel.
def half source · line 899 · raw
@n:Nat -> Nat
def leb_refl source · line 903 · raw
@+a:Nat -> {True{} == leb(a, a) : Bool}a <= a
def add_succ_r source · line 911 · raw
@+a:Nat -> @-b:Nat -> {1n+Nat.add(a, b) == Nat.add(a, 1n+b) : Nat}1 + (a + b) == a + (1 + b)
def add_zero_r source · line 920 · raw
@+a:Nat -> {Nat.add(a, 0n) == a : Nat}a + 0 == a
def half_go_le source · line 929 · raw
@+n:Nat -> @+acc:Nat -> {True{} == leb(half.go(n, acc), Nat.add(n, acc)) : Bool}half.go(n, acc) <= n + acc
def half_le source · line 942 · raw
@+n:Nat -> {True{} == leb(half(n), n) : Bool}half(n) <= n
def add_sub source · line 947 · raw
@+a:Nat -> @+r:Nat -> @e:{True{} == leb(a, r) : Bool} -> {r == Nat.add(a, Nat.sub(r, a)) : Nat}a <= r -> r == a + (r - a)
def Bands.zeros source · line 974 · raw
@+d:Nat -> @+r:Nat -> @+c:Nat -> Bands<r, c>
r x c of zeros in 2^d bands (d = 0: a single band, the same as a Mat)
def eq_half source · line 983 · raw
@-a:Nat -> @-b:Nat -> @e:{a == b : Nat} -> {half(a) == half(b) : Nat}the row counts of the two bands follow the row count of the node: if rr == r at run time, the same holds for both halves (the bands carry r only in their type)
def eq_rest source · line 987 · raw
@-a:Nat -> @-b:Nat -> @e:{a == b : Nat} -> {Nat.sub(a, half(a)) == Nat.sub(b, half(b)) : Nat}
def bfill_l source · line 993 · raw
@+r:Nat -> @+c:Nat -> @+i:Nat -> @+n:Nat -> Bool
does a block of n numbers starting at flat index i end before the second band of a node with r rows? does it start in the second band?
def bfill_r source · line 996 · raw
@+r:Nat -> @+c:Nat -> @+i:Nat -> Bool
def fill_n source · line 1000 · raw
@n:Nat -> @xs:List<&2, F32> -> @+i:U32 -> @a:Array<F32> -> Array<F32>
writes the first n numbers of xs at index i (and nothing past them)
def fill_band source · line 1007 · raw
@-r:Nat -> @-c:Nat -> @n:Nat -> @xs:List<&2, F32> -> @+i:U32 -> @m:Mat<r, c> -> Mat<r, c>
def bfill source · line 1015 · raw
@+c:Nat -> @-r:Nat -> @b:Bands<r, c> -> @l:Bool -> @rt:Bool -> @+rr:Nat -> @-e:{rr == r : Nat} -> @+i:Nat -> @+n:Nat -> @+xs:List<&2, F32> -> Bands<r, c>writes the first n numbers of xs at flat index i; l and rt say which bands they touch (bfill_l, bfill_r); rr is the row count r at run time. A block that crosses into the second band is not copied: the first band writes only its part, the second gets the rest by List.drop.
def bfill_top source · line 1026 · raw
@ok:Bool -> @+r:Nat -> @+c:Nat -> @b:Bands<r, c> -> @+i:Nat -> @+n:Nat -> @+xs:List<&2, F32> -> Bands<r, c>
def Bands.fill_at source · line 1035 · raw
@+r:Nat -> @+c:Nat -> @+i:Nat -> @+n:Nat -> @+xs:List<&2, F32> -> @b:Bands<r, c> -> Bands<r, c>
writes the n numbers of xs into b starting at flat index i (row-major, as Mat.fill_at); numbers past r*c are dropped. For loading a large matrix block by block.
def bv_g2 source · line 1048 · raw
@-r:Nat -> @-c:Nat -> @w:Array<F32> -> @al:AL -> BL<r, c>
def bv_g source · line 1054 · raw
@-r:Nat -> @-c:Nat -> @+m:Nat -> @g:G -> BL<r, c>
the band's 1 + m rows came out in an Array (one number per row); read them back as a list
def bv_gemm source · line 1062 · raw
@-r:Nat -> @+c:Nat -> @rows:Nat -> @+rr:Nat -> @x:Array<F32> -> @w:Array<F32> -> BL<r, c>
the rows of one band, one dot product each, with the gemm of Mat.matmul_nt (n = 1), so the results are identical. Writing the results into an Array and reading them back is faster than building the list row by row (NOTES.md, exp. 14).
def bv_mat source · line 1069 · raw
@-r:Nat -> @+c:Nat -> @+rr:Nat -> @x:Array<F32> -> @m:Mat<r, c> -> BL<r, c>
def bv_join source · line 1074 · raw
@-r:Nat -> @-c:Nat -> @p:BL<half(r), c> -> @q:BL<Nat.sub(r, half(r)), c> -> BL<r, c>
def bmv source · line 1083 · raw
@+c:Nat -> @-r:Nat -> @b:Bands<r, c> -> @+rr:Nat -> @-e:{rr == r : Nat} -> @xp:Pair(Array<F32>, Array<F32>) -> BL<r, c>the two bands of a node run in parallel, each with its own copy of x (c numbers); xp is a pair of copies of x (a leaf uses one). xp comes after the erased e on purpose: in Bend 2.0.35 an erased parameter at the end of the list turns the parallel let into two sequential calls (NOTES.md, exp. 11).
def bv_fin source · line 1091 · raw
@+r:Nat -> @-c:Nat -> @bl:BL<r, c> -> BV<r, c>
def Bands.matvec source · line 1098 · raw
@+r:Nat -> @+c:Nat -> @b:Bands<r, c> -> @x:Mat<1n, c> -> BV<r, c>
y = B · x for B (r x c, one row per output) and x (1 x c); returns B intact and y (1 x r). The same product as Mat.matmul_nt(1n, c, r, ..), and the same numbers, bit for bit.
def Bands.matvec_l source · line 1105 · raw
@+r:Nat -> @+c:Nat -> @b:Bands<r, c> -> @+x:List<&2, F32> -> BL<r, c>
the same product with x and y as lists (x: c numbers, y: r numbers in row order), without the Mat conversions of Bands.matvec: for callers that already hold the vector as a list
def brow_fin source · line 1113 · raw
@-r:Nat -> @-c:Nat -> @al:AL -> BRow<r, c>
def brow_leaf2 source · line 1119 · raw
@-r:Nat -> @-c:Nat -> @n:Nat -> @+i:U32 -> @d:Array<F32> -> BRow<r, c>
n numbers starting at flat index i of one band (a band may have no rows)
def brow_leaf source · line 1126 · raw
@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> @+i:U32 -> @n:Nat -> BRow<r, c>
def brow_l source · line 1131 · raw
@-r:Nat -> @-c:Nat -> @p:BRow<half(r), c> -> @y:Bands<Nat.sub(r, half(r)), c> -> BRow<r, c>
def brow_r source · line 1136 · raw
@-r:Nat -> @-c:Nat -> @x:Bands<half(r), c> -> @q:BRow<Nat.sub(r, half(r)), c> -> BRow<r, c>
def brow source · line 1142 · raw
@+c:Nat -> @-r:Nat -> @b:Bands<r, c> -> @left:Bool -> @+rr:Nat -> @-e:{rr == r : Nat} -> @+i:Nat -> BRow<r, c>row i; left says whether i is in the first band of the node (i < half(rr))
def Bands.read_row source · line 1152 · raw
@+r:Nat -> @+c:Nat -> @b:Bands<r, c> -> @+i:Nat -> BRow<r, c>
row i (c numbers), for example an embedding lookup
def bl_join source · line 1155 · raw
@-r:Nat -> @-c:Nat -> @p:BRow<half(r), c> -> @q:BRow<Nat.sub(r, half(r)), c> -> BRow<r, c>
def blist source · line 1160 · raw
@+c:Nat -> @-r:Nat -> @b:Bands<r, c> -> @+rr:Nat -> @-e:{rr == r : Nat} -> BRow<r, c>
def Bands.to_list source · line 1168 · raw
@+r:Nat -> @+c:Nat -> @b:Bands<r, c> -> BRow<r, c>
all r*c numbers in row order (in BRow.row)
def bl_len source · line 1176 · raw
@-r:Nat -> @-c:Nat -> @bl:BL<r, c> -> Nat
def succ_out source · line 1182 · raw
@-s:Nat -> @+a:Nat -> @+b:Nat -> @e:{s == Nat.add(a, 1n+b) : Nat} -> {s == 1n+Nat.add(a, b) : Nat}s == a + (1 + b) -> s == 1 + (a + b)
def len_step source · line 1187 · raw
@-s:Nat -> @+p:Nat -> @+l:Nat -> @e:{s == 1n+Nat.add(p, 1n+l) : Nat} -> {s == 2n+Nat.add(p, l) : Nat}s == 1 + (p + (1 + l)) -> s == 2 + (p + l)
def zero_out source · line 1192 · raw
@-s:Nat -> @+n:Nat -> @e:{s == Nat.add(n, 0n) : Nat} -> {s == n : Nat}s == n + 0 -> s == n
def rev_len source · line 1197 · raw
@+xs:List<&2, F32> -> @+acc:List<&2, F32> -> {List.length(&2, F32, List.reverse.go(&2, F32, xs, acc)) == Nat.add(List.length(&2, F32, xs), List.length(&2, F32, acc)) : Nat}reversing onto acc adds the lengths
def app_len source · line 1205 · raw
@+xs:List<&2, F32> -> @+ys:List<&2, F32> -> {List.length(&2, F32, List.append(&2, F32, xs, ys)) == Nat.add(List.length(&2, F32, xs), List.length(&2, F32, ys)) : Nat}appending adds the lengths
def g_len source · line 1235 · raw
@-r:Nat -> @-c:Nat -> @+p:Nat -> @g:G -> {bl_len(r, c, bv_g(r, c, p, g)) == 1n+p : Nat}whatever the gemm wrote, reading 1 + p rows back gives 1 + p numbers (G has one constructor, so matching g opens it, and leaf_len counts the list)
def join_len source · line 1261 · raw
@-r:Nat -> @-c:Nat -> @-h1:Nat -> @-h2:Nat -> @p:BL<half(r), c> -> @q:BL<Nat.sub(r, half(r)), c> -> @hp:{bl_len(half(r), c, p) == h1 : Nat} -> @hq:{bl_len(Nat.sub(r, half(r)), c, q) == h2 : Nat} -> {bl_len(r, c, bv_join(r, c, p, q)) == Nat.add(h1, h2) : Nat}joining two halves of h1 and h2 numbers gives h1 + h2 numbers. (The statement for the whole tree, "Bands.matvec_l gives r numbers", would apply this to the results of the two recursive calls and to the induction hypotheses about the same calls; that uses each band's Array twice, which Bend's affine rules refuse. It stays open, covered by tests: NOTES.md, v3.1.)