~/bend-docscommunity

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

type G source · line 27 · raw

Type

type S3 source · line 66 · raw

Type

type P2 source · line 176 · raw

Type

type AL source · line 179 · raw

Type

type L2 source · line 318 · raw

Type

type Mat source · line 331 · raw

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

type CR source · line 345 · raw

Type

type ML source · line 449 · raw

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

type MM2 source · line 464 · raw

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

type MMul source · line 481 · raw

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

type MMulNT source · line 486 · raw

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

type MMulTN source · line 491 · raw

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

type MBias source · line 560 · raw

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

type MPair source · line 575 · raw

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

type MColSum source · line 614 · raw

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

type MCE source · line 630 · raw

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

type MHits source · line 645 · raw

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

type MRow source · line 660 · raw

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

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>