~/bend-docscommunity

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

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>