oa::FnMatrix::scaledDotProductAttentionBwd

scaledDotProductAttentionBwd: explicit adjoint over provider-owned saved state. inOutput may be empty for providers whose saved state is sufficient.

Function Documentation

ScaledDotProductAttentionBwdResult oa::FnMatrix::scaledDotProductAttentionBwd( const Matrix & inQ, const Matrix & inK, const Matrix & inV, const Matrix & inOutput, const Matrix & inSavedState, const Matrix & inGradOutput, oa::F32 inScale, bool inCausal )
scaledDotProductAttentionBwd: explicit adjoint over provider-owned saved state. inOutput may be empty for providers whose saved state is sufficient.

Parameters

inQ
const Matrix &

inK
const Matrix &

inV
const Matrix &

inOutput
const Matrix &

inSavedState
const Matrix &

inGradOutput
const Matrix &

inScale
oa::F32

inCausal
bool

Returns

ScaledDotProductAttentionBwdResult

The declared return value.