oa::FnMatrix::splitHeads

splitHeads: [B*S,D] -> [B*H,S,D/H]. Multi-head permutation.

Function Documentation

Matrix oa::FnMatrix::splitHeads( const Matrix & inX, oa::I32 inBatch, oa::I32 inSeqLen, oa::I32 inNumHeads )
splitHeads: [B*S,D] -> [B*H,S,D/H]. Multi-head permutation.

Parameters

inX
const Matrix &

inBatch
oa::I32

inSeqLen
oa::I32

inNumHeads
oa::I32

Returns

Matrix

The declared return value.