shaped-matmul ( a b -- result )
NumPy-style shaped-array operations

Prev:shaped-max ( array axes keepdims? -- result )


Vocabulary
arrays.shaped

Inputs
aa shaped-array or a sequence
ba shaped-array or a sequence


Outputs
resulta shaped-array


Word description
Multiplies matrices on the final two axes and broadcasts leading batch axes. A left vector is promoted to a row, and a right vector to a column; the inserted axes are removed from the result. Two vectors produce a zero-rank dot product. Scalars and incompatible contraction dimensions are rejected. Empty contraction dimensions produce zeros. Inputs may be views; the result has fresh storage. This uses Factor arithmetic and does not conjugate complex inputs or call BLAS.

Examples
USING: arrays.shaped accessors prettyprint ; { 1 2 3 } { 4 5 6 } shaped-matmul underlying>> .
{ 32 }


Definition


:: shaped-matmul ( a b -- result )
a >shaped-array :> left b >shaped-array :> right left
shape>> :> ls right shape>> :> rs ls empty? rs empty? or
[ ls rs invalid-shaped-matmul ] when ls length 1 =
:> left-vector? rs length 1 = :> right-vector? left-vector?
[ ls 1 prefix ] [ ls ] if :> lm right-vector?
[ rs 1 suffix ] [ rs ] if
:> rm lm last :> contracted rm length 2 - rm nth
contracted = [ ls rs invalid-shaped-matmul ] unless
lm lm length 2 - head rm rm length 2 - head
matmul-batch-shape :> batch lm length 2 - lm nth
:> rows rm last :> columns batch rows suffix columns suffix
:> full-shape lm full-shape length broadcast-strides
:> left-strides rm full-shape length broadcast-strides
:> right-strides full-shape product <iota> [| index |
index full-shape flat-coordinate
:> coordinate coordinate batch length head
:> batch-coordinate batch length coordinate nth
:> row coordinate last :> column contracted <iota> [| k
|
batch-coordinate row suffix k suffix
left-strides coordinate-offset left underlying>> nth
batch-coordinate k suffix column suffix
right-strides coordinate-offset
right underlying>> nth *
] map sum
] map batch left-vector? ~8 more~ ;