oa::FnMatrix::rmsNormGatedBwd

rmsNormGatedBwd: backward for rmsNormGated (normBeforeGate = true). Returns grads w.r.t. x, weight, bias, z.

Function Documentation

RmsNormGatedBwdResult oa::FnMatrix::rmsNormGatedBwd( const Matrix & inX, const Matrix & inWeight, const Matrix & inBias, const Matrix & inZ, const Matrix & inGradOutput, oa::F32 inEps )
rmsNormGatedBwd: backward for rmsNormGated (normBeforeGate = true). Returns grads w.r.t. x, weight, bias, z.

Parameters

inX
const Matrix &

inWeight
const Matrix &

inBias
const Matrix &

inZ
const Matrix &

inGradOutput
const Matrix &

inEps
oa::F32

Returns

RmsNormGatedBwdResult

The declared return value.