oa::MultiHeadAttention

MultiHeadAttention: Multi-head scaled dot-product self-attention with explicit causal/bidirectional visibility and interchangeable standard/fused causal Flash backends.

Inheritance

public Module

Public Methods

oa::AttentionBackend oa::MultiHeadAttention::backend()
Matrix oa::MultiHeadAttention::forward(const Matrix & inInput)
oa::Matrix oa::MultiHeadAttention::forwardMasked(const oa::Matrix & inInput, const oa::Matrix & inAdditiveMask)
oa::I32 oa::MultiHeadAttention::headDim()
oa::AttentionBackend oa::MultiHeadAttention::lastBackend()
oa::AttentionMode oa::MultiHeadAttention::mode()
oa::I32 oa::MultiHeadAttention::numHeads()
void oa::MultiHeadAttention::setBackend(oa::AttentionBackend inBackend)
void oa::MultiHeadAttention::setMode(oa::AttentionMode inMode)
void oa::MultiHeadAttention::setSeqLen(oa::I32 inSeqLen)

Constructor & Destructor Documentation

oa::MultiHeadAttention::MultiHeadAttention( oa::I32 inDModel, oa::I32 inNumHeads, oa::F32 inDropoutP = 0.0f, bool inBias = true, oa::AttentionBackend inBackend = oa::AttentionBackend::Auto, oa::AttentionMode inMode = oa::AttentionMode::Causal )
No public source comment is attached to this declaration.

Parameters

inDModel
oa::I32

inNumHeads
oa::I32

inDropoutP
oa::F32

Default: 0.0f

inBias
bool

Default: true

inBackend
oa::AttentionBackend

Default: oa::AttentionBackend::Auto

inMode
oa::AttentionMode

Default: oa::AttentionMode::Causal

Public Method Documentation

oa::AttentionBackend oa::MultiHeadAttention::backend()
No public source comment is attached to this declaration.

Returns

oa::AttentionBackend

The declared return value.

Matrix oa::MultiHeadAttention::forward( const Matrix & inInput )
No public source comment is attached to this declaration.

Parameters

inInput
const Matrix &

Returns

Matrix

The declared return value.

oa::Matrix oa::MultiHeadAttention::forwardMasked( const oa::Matrix & inInput, const oa::Matrix & inAdditiveMask )
No public source comment is attached to this declaration.

Parameters

inInput
const oa::Matrix &

inAdditiveMask
const oa::Matrix &

Returns

oa::Matrix

The declared return value.

oa::I32 oa::MultiHeadAttention::headDim()
No public source comment is attached to this declaration.

Returns

oa::I32

The declared return value.

oa::AttentionBackend oa::MultiHeadAttention::lastBackend()
No public source comment is attached to this declaration.

Returns

oa::AttentionBackend

The declared return value.

oa::AttentionMode oa::MultiHeadAttention::mode()
No public source comment is attached to this declaration.

Returns

oa::AttentionMode

The declared return value.

oa::I32 oa::MultiHeadAttention::numHeads()
No public source comment is attached to this declaration.

Returns

oa::I32

The declared return value.

void oa::MultiHeadAttention::setBackend( oa::AttentionBackend inBackend )
No public source comment is attached to this declaration.

Parameters

inBackend
oa::AttentionBackend

Returns

void

The declared return value.

void oa::MultiHeadAttention::setMode( oa::AttentionMode inMode )
No public source comment is attached to this declaration.

Parameters

inMode
oa::AttentionMode

Returns

void

The declared return value.

void oa::MultiHeadAttention::setSeqLen( oa::I32 inSeqLen )
No public source comment is attached to this declaration.

Parameters

inSeqLen
oa::I32

Returns

void

The declared return value.