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