main.bend checks
raw source on the hub · import bend-ml-tensor-array@0.1.0.0/main.bend as Main
bend-ml-tensor-array: tensores sobre Array<F32> plano, com a shape no tipo.
import bend-ml-tensor-array@0.1.0.0/main.bend as TA
Mat<r, c> guarda r*c números num Array<F32> linha a linha (índice i*c + j). As dimensões são parâmetros de tipo apagados: a mesma garantia do bend-ml-tensor (produto com dimensões erradas não compila), mas ~50x mais rápido que listas: Array.get/set por índice, e produtos em blocos de linhas em paralelo.
Um Array é afim (um só dono): toda operação que lê um operando o devolve junto com o resultado (como Array.get faz), em registros MM, AR, ... Use Mat.clone quando precisar de duas cópias.
O que NÃO é provado: a invariante "o Array tem capacidade >= r*c" vale porque os construtores (Mat.zeros, Mat.of_list) alocam o tamanho certo; e F32 é validado por testes contra o PyTorch (reference/test_tensor_array.py), não por prova.
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 135 · raw
Type
P2@a:Array<F32> -> @b:Array<F32> -> P2
type AL source · line 138 · raw
Type
AL@c:Array<F32> -> @l:List<&2, F32> -> AL
type L2 source · line 277 · raw
Type
L2@loss:F32 -> @c:Array<F32> -> L2
type Mat source · line 290 · raw
@-r:Nat -> @-c:Nat -> Type
Mat@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>
type CR source · line 304 · raw
Type
CR@c:Array<F32> -> @n:Nat -> CR
type ML source · line 402 · 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 417 · 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 434 · 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 439 · 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 444 · 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 503 · 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 518 · 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 557 · 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 573 · raw
@-n:Nat -> @-c:Nat -> Type
MCE@-n:Nat -> @-c:Nat -> @loss:F32 -> @grad:Mat<n, c> -> MCE<n, c>
type MHits source · line 588 · raw
@-n:Nat -> @-c:Nat -> Type
MHits@-n:Nat -> @-c:Nat -> @z:Mat<n, c> -> @hits:Nat -> MHits<n, c>
type MRow source · line 603 · 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 depth source · line 125 · raw
@n:Nat -> Nat
def read_l source · line 143 · raw
@mleft:Nat -> @+idx:U32 -> @acc:List<&2, F32> -> @rc:Pair(Array<F32>, F32) -> AL
def addl source · line 152 · raw
@xs:List<&2, F32> -> @+idx:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>
def bias_rows source · line 161 · raw
@nleft:Nat -> @+i:U32 -> @+mu:U32 -> @+bl:List<&2, F32> -> @c:Array<F32> -> Array<F32>
def relu_l source · line 170 · raw
@nleft:Nat -> @+t:U32 -> @rc:Pair(Array<F32>, F32) -> Array<F32>
def pos2 source · line 177 · raw
@b:Bool -> F32
def pos source · line 184 · raw
@+y:F32 -> F32
def mask_l source · line 189 · raw
@nleft:Nat -> @+t:U32 -> @rd:Pair(Array<F32>, F32) -> @rh:Pair(Array<F32>, F32) -> P2
def maxl2 source · line 198 · raw
@xs:List<&2, F32> -> @cur:F32 -> F32
def maxl source · line 205 · raw
@xs:List<&2, F32> -> F32
def expl source · line 212 · raw
@xs:List<&2, F32> -> @+m:F32 -> List<&2, F32>
def suml source · line 219 · raw
@xs:List<&2, F32> -> @acc:F32 -> F32
def divl source · line 226 · raw
@xs:List<&2, F32> -> @+s:F32 -> List<&2, F32>
def softmax2 source · line 233 · raw
@+es:List<&2, F32> -> List<&2, F32>
def softmax source · line 236 · raw
@+xs:List<&2, F32> -> List<&2, F32>
def sgd_l source · line 241 · raw
@nleft:Nat -> @+t:U32 -> @+lr:F32 -> @rw:Pair(Array<F32>, F32) -> @rd:Pair(Array<F32>, F32) -> P2
def pick source · line 252 · raw
@ps:List<&2, F32> -> @n:Nat -> F32
def hit source · line 263 · raw
@b:Bool -> F32
def grad_l source · line 270 · raw
@ps:List<&2, F32> -> @+label:Nat -> @+j:Nat -> @+nf:F32 -> List<&2, F32>
def ce_row2 source · line 280 · raw
@+ps:List<&2, F32> -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> @c:Array<F32> -> L2
def ce_row source · line 283 · raw
@al:AL -> @+label:Nat -> @+nf:F32 -> @+idx:U32 -> @loss:F32 -> L2
def ce_rows source · line 295 · raw
@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+nf:F32 -> @+i:U32 -> @r:L2 -> L2
def gt_head source · line 307 · raw
@xs:List<&2, F32> -> @+v:F32 -> Bool
def argmax2 source · line 314 · raw
@xs:List<&2, F32> -> @g:Bool -> @+i:Nat -> @+bi:Nat -> @+bv:F32 -> Nat
def argmax source · line 323 · raw
@xs:List<&2, F32> -> Nat
def one source · line 330 · raw
@b:Bool -> Nat
def count_rows3 source · line 337 · raw
@+y:Nat -> @n:Nat -> @al:AL -> CR
def count_rows source · line 342 · raw
@labels:List<&2, Nat> -> @+mm1:Nat -> @+mu:U32 -> @+i:U32 -> @r:CR -> CR
def cap_depth2 source · line 354 · raw
@small:Bool -> @n:Nat -> Nat
def cap_depth source · line 361 · raw
@+n:Nat -> Nat
def Mat.zeros source · line 366 · raw
@+r:Nat -> @+c:Nat -> Mat<r, c>
def Mat.fill source · line 369 · raw
@+r:Nat -> @+c:Nat -> @+v:F32 -> Mat<r, c>
def fill_list source · line 372 · raw
@xs:List<&2, F32> -> @+i:U32 -> @a:Array<F32> -> Array<F32>
def len_is source · line 379 · raw
@xs:List<&2, F32> -> @n:Nat -> Bool
def Mat.of_list2 source · line 388 · 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 397 · raw
@+r:Nat -> @+c:Nat -> @+xs:List<&2, F32> -> Maybe<&1, Mat<r, c>>
def Mat.to_list2 source · line 405 · raw
@-r:Nat -> @-c:Nat -> @al:AL -> ML<r, c>
def Mat.to_list source · line 412 · raw
@+r:Nat -> @+c:Nat -> @m:Mat<r, c> -> ML<r, c>
def Mat.clone2 source · line 420 · raw
@-r:Nat -> @-c:Nat -> @p:Pair(Array<F32>, Array<F32>) -> MM2<r, c>
def Mat.clone source · line 429 · raw
@-r:Nat -> @-c:Nat -> @m:Mat<r, c> -> MM2<r, c>
def gemm_any source · line 447 · 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 454 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMul<n, k, m>
def mmul_nt_fin source · line 459 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulNT<n, k, m>
def mmul_tn_fin source · line 464 · raw
@-n:Nat -> @-k:Nat -> @-m:Nat -> @g:G -> MMulTN<n, k, m>
def mm_a source · line 469 · raw
@+n:Nat -> @+k:Nat -> @+m:Nat -> @d:Nat -> @a:Array<F32> -> @b:Array<F32> -> MMul<n, k, m>
def Mat.matmul source · line 474 · 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 479 · 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 484 · 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 489 · 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 494 · 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 506 · raw
@-n:Nat -> @-m:Nat -> @+nn:Nat -> @+mm:Nat -> @cd:Array<F32> -> @al:AL -> MBias<n, m>
def Mat.add_row source · line 511 · raw
@+n:Nat -> @+m:Nat -> @c:Mat<n, m> -> @b:Mat<1n, m> -> MBias<n, m>
def mpair_fin source · line 521 · raw
@-r:Nat -> @-c:Nat -> @p:P2 -> MPair<r, c>
def relu_fin source · line 526 · raw
@-r:Nat -> @-c:Nat -> @d:Array<F32> -> Mat<r, c>
def Mat.relu source · line 529 · raw
@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> Mat<r, c>
def Mat.relu_bwd source · line 536 · raw
@+r:Nat -> @+c:Nat -> @dy:Mat<r, c> -> @h:Mat<r, c> -> MPair<r, c>
def Mat.sgd source · line 543 · raw
@+r:Nat -> @+c:Nat -> @+lr:F32 -> @w:Mat<r, c> -> @dw:Mat<r, c> -> MPair<r, c>
def Mat.add source · line 550 · raw
@+r:Nat -> @+c:Nat -> @a:Mat<r, c> -> @b:Mat<r, c> -> MPair<r, c>
def colsum_fin source · line 560 · raw
@-n:Nat -> @-m:Nat -> @g:G -> MColSum<n, m>
def Mat.col_sums source · line 565 · raw
@+n:Nat -> @+m:Nat -> @a:Mat<n, m> -> MColSum<n, m>
def ce_fin source · line 576 · raw
@-n:Nat -> @-c:Nat -> @+nf:F32 -> @r:L2 -> MCE<n, c>
def Mat.softmax_ce source · line 581 · raw
@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MCE<n, c>
def hits_fin source · line 591 · raw
@-n:Nat -> @-c:Nat -> @r:CR -> MHits<n, c>
def Mat.count_correct source · line 596 · raw
@+n:Nat -> @+c:Nat -> @z:Mat<n, c> -> @labels:List<&2, Nat> -> MHits<n, c>
def row_fin source · line 606 · raw
@-r:Nat -> @-c:Nat -> @al:AL -> MRow<r, c>
def Mat.read_row source · line 611 · raw
@+r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> MRow<r, c>
def Mat.write_row source · line 616 · raw
@-r:Nat -> @+c:Nat -> @m:Mat<r, c> -> @+i:U32 -> @xs:List<&2, F32> -> Mat<r, c>