30#ifndef TMVA_DNN_GRU_LAYER
31#define TMVA_DNN_GRU_LAYER
57template<
typename Architecture_t>
63 using Matrix_t =
typename Architecture_t::Matrix_t;
64 using Scalar_t =
typename Architecture_t::Scalar_t;
65 using Tensor_t =
typename Architecture_t::Tensor_t;
193 void Print()
const override;
309template <
typename Architecture_t>
321 fCandidateValue(batchSize,
stateSize), fState(batchSize,
stateSize), fWeightsResetGate(
this->GetWeightsAt(0)),
322 fWeightsResetGateState(
this->GetWeightsAt(3)), fResetGateBias(
this->GetBiasesAt(0)),
323 fWeightsUpdateGate(
this->GetWeightsAt(1)), fWeightsUpdateGateState(
this->GetWeightsAt(4)),
324 fUpdateGateBias(
this->GetBiasesAt(1)), fWeightsCandidate(
this->GetWeightsAt(2)),
325 fWeightsCandidateState(
this->GetWeightsAt(5)), fCandidateBias(
this->GetBiasesAt(2)),
326 fWeightsResetGradients(
this->GetWeightGradientsAt(0)), fWeightsResetStateGradients(
this->GetWeightGradientsAt(3)),
327 fResetBiasGradients(
this->GetBiasGradientsAt(0)), fWeightsUpdateGradients(
this->GetWeightGradientsAt(1)),
328 fWeightsUpdateStateGradients(
this->GetWeightGradientsAt(4)), fUpdateBiasGradients(
this->GetBiasGradientsAt(1)),
329 fWeightsCandidateGradients(
this->GetWeightGradientsAt(2)),
330 fWeightsCandidateStateGradients(
this->GetWeightGradientsAt(5)),
331 fCandidateBiasGradients(
this->GetBiasGradientsAt(2))
341 Architecture_t::InitializeGRUTensors(
this);
345template <
typename Architecture_t>
348 fStateSize(
layer.fStateSize),
349 fTimeSteps(
layer.fTimeSteps),
350 fRememberState(
layer.fRememberState),
351 fReturnSequence(
layer.fReturnSequence),
352 fResetGateAfter(
layer.fResetGateAfter),
353 fF1(
layer.GetActivationFunctionF1()),
354 fF2(
layer.GetActivationFunctionF2()),
355 fResetValue(
layer.GetBatchSize(),
layer.GetStateSize()),
356 fUpdateValue(
layer.GetBatchSize(),
layer.GetStateSize()),
357 fCandidateValue(
layer.GetBatchSize(),
layer.GetStateSize()),
358 fState(
layer.GetBatchSize(),
layer.GetStateSize()),
359 fWeightsResetGate(
this->GetWeightsAt(0)),
360 fWeightsResetGateState(
this->GetWeightsAt(3)),
361 fResetGateBias(
this->GetBiasesAt(0)),
362 fWeightsUpdateGate(
this->GetWeightsAt(1)),
363 fWeightsUpdateGateState(
this->GetWeightsAt(4)),
364 fUpdateGateBias(
this->GetBiasesAt(1)),
365 fWeightsCandidate(
this->GetWeightsAt(2)),
366 fWeightsCandidateState(
this->GetWeightsAt(5)),
367 fCandidateBias(
this->GetBiasesAt(2)),
368 fWeightsResetGradients(
this->GetWeightGradientsAt(0)),
369 fWeightsResetStateGradients(
this->GetWeightGradientsAt(3)),
370 fResetBiasGradients(
this->GetBiasGradientsAt(0)),
371 fWeightsUpdateGradients(
this->GetWeightGradientsAt(1)),
372 fWeightsUpdateStateGradients(
this->GetWeightGradientsAt(4)),
373 fUpdateBiasGradients(
this->GetBiasGradientsAt(1)),
374 fWeightsCandidateGradients(
this->GetWeightGradientsAt(2)),
375 fWeightsCandidateStateGradients(
this->GetWeightGradientsAt(5)),
376 fCandidateBiasGradients(
this->GetBiasGradientsAt(2))
406 Architecture_t::InitializeGRUTensors(
this);
410template <
typename Architecture_t>
415 Architecture_t::InitializeGRUDescriptors(fDescriptors,
this);
416 Architecture_t::InitializeGRUWorkspace(fWorkspace, fDescriptors,
this);
419 if (Architecture_t::IsCudnn())
420 fResetGateAfter =
true;
424template <
typename Architecture_t>
433 Architecture_t::MultiplyTranspose(
tmpState, fState, fWeightsResetGateState);
434 Architecture_t::MultiplyTranspose(fResetValue,
input, fWeightsResetGate);
435 Architecture_t::ScaleAdd(fResetValue,
tmpState);
436 Architecture_t::AddRowWise(fResetValue, fResetGateBias);
437 DNN::evaluateDerivativeMatrix<Architecture_t>(
dr,
fRst, fResetValue);
438 DNN::evaluateMatrix<Architecture_t>(fResetValue,
fRst);
442template <
typename Architecture_t>
451 Architecture_t::MultiplyTranspose(
tmpState, fState, fWeightsUpdateGateState);
452 Architecture_t::MultiplyTranspose(fUpdateValue,
input, fWeightsUpdateGate);
453 Architecture_t::ScaleAdd(fUpdateValue,
tmpState);
454 Architecture_t::AddRowWise(fUpdateValue, fUpdateGateBias);
455 DNN::evaluateDerivativeMatrix<Architecture_t>(
du,
fUpd, fUpdateValue);
456 DNN::evaluateMatrix<Architecture_t>(fUpdateValue,
fUpd);
460template <
typename Architecture_t>
477 Matrix_t tmp(fCandidateValue.GetNrows(), fCandidateValue.GetNcols());
478 if (!fResetGateAfter) {
480 Architecture_t::Hadamard(
tmpState, fState);
481 Architecture_t::MultiplyTranspose(
tmp,
tmpState, fWeightsCandidateState);
484 Architecture_t::MultiplyTranspose(
tmp, fState, fWeightsCandidateState);
485 Architecture_t::Hadamard(
tmp, fResetValue);
487 Architecture_t::MultiplyTranspose(fCandidateValue,
input, fWeightsCandidate);
488 Architecture_t::ScaleAdd(fCandidateValue,
tmp);
489 Architecture_t::AddRowWise(fCandidateValue, fCandidateBias);
490 DNN::evaluateDerivativeMatrix<Architecture_t>(
dc, fCan, fCandidateValue);
491 DNN::evaluateMatrix<Architecture_t>(fCandidateValue, fCan);
495template <
typename Architecture_t>
500 if (Architecture_t::IsCudnn()) {
507 Architecture_t::Rearrange(
x,
input);
510 const auto &weights = this->GetWeightsTensor();
512 auto &
hx = this->fState;
513 auto &
cx = this->fCell;
515 auto &
hy = this->fState;
516 auto &
cy = this->fCell;
523 if (fReturnSequence) {
524 Architecture_t::Rearrange(this->GetOutput(),
y);
527 Tensor_t tmp = (
y.At(
y.GetShape()[0] - 1)).Reshape({
y.GetShape()[1], 1,
y.GetShape()[2]});
528 Architecture_t::Copy(this->GetOutput(),
tmp);
539 Tensor_t arrInput ( fTimeSteps, this->GetBatchSize(), this->GetInputWidth());
550 if (!this->fRememberState) {
556 for (
size_t t = 0; t < fTimeSteps; ++t) {
558 ResetGate(
arrInput[t], fDerivativesReset[t]);
559 Architecture_t::Copy(this->GetResetGateTensorAt(t), fResetValue);
560 UpdateGate(
arrInput[t], fDerivativesUpdate[t]);
561 Architecture_t::Copy(this->GetUpdateGateTensorAt(t), fUpdateValue);
563 CandidateValue(
arrInput[t], fDerivativesCandidate[t]);
564 Architecture_t::Copy(this->GetCandidateGateTensorAt(t), fCandidateValue);
567 CellForward(fUpdateValue, fCandidateValue);
576 Architecture_t::Rearrange(this->GetOutput(),
arrOutput);
582 tmp =
tmp.Reshape({
tmp.GetShape()[0],
tmp.GetShape()[1], 1});
584 assert(
tmp.GetShape()[0] ==
this->GetOutput().GetShape()[2]);
585 Architecture_t::Rearrange(this->GetOutput(),
tmp);
592template <
typename Architecture_t>
600 for (
size_t j = 0;
j < (size_t)
tmp.GetNcols();
j++) {
601 for (
size_t i = 0; i < (size_t)
tmp.GetNrows(); i++) {
612template <
typename Architecture_t>
618 if (Architecture_t::IsCudnn()) {
631 if (!fReturnSequence) {
634 Architecture_t::InitializeZero(
dy);
640 Architecture_t::Copy(
tmp2, this->GetActivationGradients());
642 Architecture_t::Rearrange(
y, this->GetOutput());
643 Architecture_t::Rearrange(
dy, this->GetActivationGradients());
649 const auto &weights = this->GetWeightsTensor();
657 auto &
hx = this->GetState();
658 auto &
cx = this->GetCell();
668 Architecture_t::RNNBackward(
x,
hx,
cx,
y,
dy,
dhy,
dcy, weights,
dx,
dhx,
dcx,
weightGradients,
rnnDesc,
rnnWork);
709 if (fReturnSequence) {
710 Architecture_t::Rearrange(
arr_output, this->GetOutput());
711 Architecture_t::Rearrange(
arr_actgradients, this->GetActivationGradients());
720 this->GetActivationGradients().GetShape()[2]);
722 Architecture_t::Rearrange(
tmp_grad, this->GetActivationGradients());
729 fWeightsResetGradients.Zero();
730 fWeightsResetStateGradients.Zero();
731 fResetBiasGradients.Zero();
734 fWeightsUpdateGradients.Zero();
735 fWeightsUpdateStateGradients.Zero();
736 fUpdateBiasGradients.Zero();
739 fWeightsCandidateGradients.Zero();
740 fWeightsCandidateStateGradients.Zero();
741 fCandidateBiasGradients.Zero();
744 for (
size_t t = fTimeSteps; t > 0; t--) {
752 this->GetResetGateTensorAt(t-1), this->GetUpdateGateTensorAt(t-1),
753 this->GetCandidateGateTensorAt(t-1),
755 fDerivativesReset[t-1], fDerivativesUpdate[t-1],
756 fDerivativesCandidate[t-1]);
761 this->GetResetGateTensorAt(t-1), this->GetUpdateGateTensorAt(t-1),
762 this->GetCandidateGateTensorAt(t-1),
764 fDerivativesReset[t-1], fDerivativesUpdate[t-1],
765 fDerivativesCandidate[t-1]);
777template <
typename Architecture_t>
789 fWeightsResetGradients, fWeightsUpdateGradients, fWeightsCandidateGradients,
790 fWeightsResetStateGradients, fWeightsUpdateStateGradients,
791 fWeightsCandidateStateGradients, fResetBiasGradients, fUpdateBiasGradients,
792 fCandidateBiasGradients,
dr,
du,
dc,
795 fWeightsResetGate, fWeightsUpdateGate, fWeightsCandidate,
796 fWeightsResetGateState, fWeightsUpdateGateState, fWeightsCandidateState,
802template <
typename Architecture_t>
810template<
typename Architecture_t>
814 std::cout <<
" GRU Layer: \t ";
815 std::cout <<
" (NInput = " << this->GetInputSize();
816 std::cout <<
", NState = " << this->GetStateSize();
817 std::cout <<
", NTime = " << this->GetTimeSteps() <<
" )";
818 std::cout <<
"\tOutput = ( " << this->GetOutput().GetFirstSize() <<
" , " << this->GetOutput()[0].GetNrows() <<
" , " << this->GetOutput()[0].GetNcols() <<
" )\n";
822template <
typename Architecture_t>
837 this->WriteMatrixToXML(
layerxml,
"ResetWeights", this->GetWeightsAt(0));
838 this->WriteMatrixToXML(
layerxml,
"ResetStateWeights", this->GetWeightsAt(1));
839 this->WriteMatrixToXML(
layerxml,
"ResetBiases", this->GetBiasesAt(0));
840 this->WriteMatrixToXML(
layerxml,
"UpdateWeights", this->GetWeightsAt(2));
841 this->WriteMatrixToXML(
layerxml,
"UpdateStateWeights", this->GetWeightsAt(3));
842 this->WriteMatrixToXML(
layerxml,
"UpdateBiases", this->GetBiasesAt(1));
843 this->WriteMatrixToXML(
layerxml,
"CandidateWeights", this->GetWeightsAt(4));
844 this->WriteMatrixToXML(
layerxml,
"CandidateStateWeights", this->GetWeightsAt(5));
845 this->WriteMatrixToXML(
layerxml,
"CandidateBiases", this->GetBiasesAt(2));
849template <
typename Architecture_t>
854 this->ReadMatrixXML(parent,
"ResetWeights", this->GetWeightsAt(0));
855 this->ReadMatrixXML(parent,
"ResetStateWeights", this->GetWeightsAt(1));
856 this->ReadMatrixXML(parent,
"ResetBiases", this->GetBiasesAt(0));
857 this->ReadMatrixXML(parent,
"UpdateWeights", this->GetWeightsAt(2));
858 this->ReadMatrixXML(parent,
"UpdateStateWeights", this->GetWeightsAt(3));
859 this->ReadMatrixXML(parent,
"UpdateBiases", this->GetBiasesAt(1));
860 this->ReadMatrixXML(parent,
"CandidateWeights", this->GetWeightsAt(4));
861 this->ReadMatrixXML(parent,
"CandidateStateWeights", this->GetWeightsAt(5));
862 this->ReadMatrixXML(parent,
"CandidateBiases", this->GetBiasesAt(2));
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void input
const Matrix_t & GetWeightsCandidate() const
Matrix_t & GetWeightsCandidateStateGradients()
typename Architecture_t::RecurrentDescriptor_t LayerDescriptor_t
void Forward(Tensor_t &input, bool isTraining=true) override
Computes the next hidden state and next cell state with given input matrix.
Matrix_t & GetWeightsResetGate()
Matrix_t & fResetBiasGradients
Gradients w.r.t the reset gate - bias weights.
std::vector< Matrix_t > & GetUpdateGateTensor()
typename Architecture_t::Tensor_t Tensor_t
std::vector< Matrix_t > reset_gate_value
Reset gate value for every time step.
Matrix_t & CellBackward(Matrix_t &state_gradients_backward, const Matrix_t &precStateActivations, const Matrix_t &reset_gate, const Matrix_t &update_gate, const Matrix_t &candidate_gate, const Matrix_t &input, Matrix_t &input_gradient, Matrix_t &dr, Matrix_t &du, Matrix_t &dc)
Backward for a single time unit a the corresponding call to Forward(...).
size_t fStateSize
Hidden state size for GRU.
const Matrix_t & GetWeightsResetGradients() const
const Matrix_t & GetUpdateBiasGradients() const
bool fReturnSequence
Return in output full sequence or just last element.
const Matrix_t & GetWeightsResetStateGradients() const
std::vector< Matrix_t > fDerivativesReset
First fDerivatives of the activations reset gate.
const Tensor_t & GetWeightsTensor() const
std::vector< Matrix_t > & GetResetGateTensor()
Matrix_t & GetWeightsUpdateGateState()
const std::vector< Matrix_t > & GetCandidateGateTensor() const
const Matrix_t & GetUpdateDerivativesAt(size_t i) const
Matrix_t & GetWeightsUpdateStateGradients()
void Print() const override
Prints the info about the layer.
size_t GetInputSize() const
Getters.
Matrix_t fState
Hidden state of GRU.
Matrix_t & GetWeightsResetGradients()
Tensor_t & GetWeightGradientsTensor()
const Matrix_t & GetCandidateBias() const
std::vector< Matrix_t > update_gate_value
Update gate value for every time step.
Tensor_t & GetWeightsTensor()
Tensor_t fX
cached input tensor as T x B x I
Matrix_t & GetCandidateGateTensorAt(size_t i)
Matrix_t & GetResetBiasGradients()
Matrix_t & GetCandidateValue()
void AddWeightsXMLTo(void *parent) override
Writes the information and the weights about the layer in an XML node.
Matrix_t & GetWeightsResetGateState()
DNN::EActivationFunction fF1
Activation function: sigmoid.
const Matrix_t & GetWeightsUpdateGate() const
const std::vector< Matrix_t > & GetDerivativesReset() const
const Matrix_t & GetUpdateGateBias() const
Matrix_t & fWeightsResetGradients
Gradients w.r.t the reset gate - input weights.
std::vector< Matrix_t > & GetDerivativesUpdate()
Matrix_t & fCandidateBiasGradients
Gradients w.r.t the candidate gate - bias weights.
Matrix_t & fCandidateBias
Candidate Gate bias.
Matrix_t & GetUpdateGateTensorAt(size_t i)
DNN::EActivationFunction fF2
Activation function: tanh.
const Matrix_t & GetWeightsUpdateGradients() const
Matrix_t & GetWeightsCandidateGradients()
Matrix_t & fWeightsUpdateStateGradients
Gradients w.r.t the update gate - hidden state weights.
Matrix_t & GetWeightsCandidate()
Matrix_t & fWeightsUpdateGradients
Gradients w.r.t the update gate - input weights.
Matrix_t & GetUpdateGateValue()
size_t fTimeSteps
Timesteps for GRU.
std::vector< Matrix_t > fDerivativesCandidate
First fDerivatives of the activations candidate gate.
const Tensor_t & GetWeightGradientsTensor() const
typename Architecture_t::FilterDescriptor_t WeightsDescriptor_t
Tensor_t fWeightGradientsTensor
Tensor for all weight gradients.
Matrix_t & fUpdateBiasGradients
Gradients w.r.t the update gate - bias weights.
Matrix_t & GetWeightsResetStateGradients()
std::vector< Matrix_t > & GetCandidateGateTensor()
Matrix_t & fWeightsResetGate
Reset Gate weights for input, fWeights[0].
const Matrix_t & GetResetDerivativesAt(size_t i) const
Matrix_t & GetWeightsUpdateGate()
typename Architecture_t::Matrix_t Matrix_t
const Matrix_t & GetCandidateGateTensorAt(size_t i) const
Matrix_t & GetWeightsCandidateState()
const Matrix_t & GetCandidateBiasGradients() const
Matrix_t & GetResetGateTensorAt(size_t i)
Matrix_t & fResetGateBias
Input Gate bias.
const std::vector< Matrix_t > & GetResetGateTensor() const
Matrix_t fCell
Empty matrix for GRU.
std::vector< Matrix_t > candidate_gate_value
Candidate gate value for every time step.
typename Architecture_t::Scalar_t Scalar_t
const Matrix_t & GetWeigthsUpdateStateGradients() const
const Matrix_t & GetCandidateValue() const
Matrix_t & GetCandidateBiasGradients()
Matrix_t & fWeightsCandidateStateGradients
Gradients w.r.t the candidate gate - hidden state weights.
const std::vector< Matrix_t > & GetDerivativesUpdate() const
const Matrix_t & GetCell() const
void Initialize() override
Initialize the weights according to the given initialization method.
void UpdateGate(const Matrix_t &input, Matrix_t &df)
Forgets the past values (NN with Sigmoid)
const Matrix_t & GetCandidateDerivativesAt(size_t i) const
Matrix_t fResetValue
Computed reset gate values.
DNN::EActivationFunction GetActivationFunctionF2() const
Matrix_t & GetResetGateBias()
typename Architecture_t::RNNWorkspace_t RNNWorkspace_t
Matrix_t fUpdateValue
Computed forget gate values.
const Matrix_t & GetResetBiasGradients() const
bool fResetGateAfter
GRU variant to Apply the reset gate multiplication afterwards (used by cuDNN)
const Matrix_t & GetWeightsCandidateGradients() const
DNN::EActivationFunction GetActivationFunctionF1() const
bool DoesReturnSequence() const
Matrix_t & GetUpdateBiasGradients()
const Matrix_t & GetUpdateGateTensorAt(size_t i) const
Matrix_t & fWeightsResetGateState
Input Gate weights for prev state, fWeights[1].
Matrix_t & fWeightsUpdateGateState
Update Gate weights for prev state, fWeights[3].
const std::vector< Matrix_t > & GetDerivativesCandidate() const
Tensor_t fWeightsTensor
Tensor for all weights.
typename Architecture_t::RNNDescriptors_t RNNDescriptors_t
const Matrix_t & GetResetGateBias() const
Matrix_t & GetResetDerivativesAt(size_t i)
const Matrix_t & GetUpdateGateValue() const
const Matrix_t & GetResetGateTensorAt(size_t i) const
TDescriptors * fDescriptors
Keeps all the RNN descriptors.
void CellForward(Matrix_t &updateGateValues, Matrix_t &candidateValues)
Forward for a single cell (time unit)
Matrix_t & GetWeightsUpdateGradients()
Matrix_t & fWeightsResetStateGradients
Gradients w.r.t the reset gate - hidden state weights.
Matrix_t & fWeightsCandidateState
Candidate Gate weights for prev state, fWeights[5].
size_t GetStateSize() const
std::vector< Matrix_t > & GetDerivativesReset()
Matrix_t & fUpdateGateBias
Update Gate bias.
void Backward(Tensor_t &gradients_backward, const Tensor_t &activations_backward) override
Backpropagates the error.
const Matrix_t & GetWeightsCandidateStateGradients() const
void ResetGate(const Matrix_t &input, Matrix_t &di)
Decides the values we'll update (NN with Sigmoid)
const Matrix_t & GetWeightsResetGate() const
Tensor_t fDx
cached gradient on the input (output of backward) as T x B x I
Matrix_t & GetCandidateBias()
typename Architecture_t::TensorDescriptor_t TensorDescriptor_t
bool fRememberState
Remember state in next pass.
Matrix_t & fWeightsCandidate
Candidate Gate weights for input, fWeights[4].
Matrix_t & fWeightsCandidateGradients
Gradients w.r.t the candidate gate - input weights.
const Matrix_t & GetWeightsCandidateState() const
void ReadWeightsFromXML(void *parent) override
Read the information and the weights about the layer from XML node.
const std::vector< Matrix_t > & GetUpdateGateTensor() const
const Matrix_t & GetResetGateValue() const
void Update(const Scalar_t learningRate)
Tensor_t fY
cached output tensor as T x B x S
Matrix_t fCandidateValue
Computed candidate values.
const Matrix_t & GetState() const
Matrix_t & GetUpdateGateBias()
void InitState(DNN::EInitialization m=DNN::EInitialization::kZero)
Initialize the hidden state and cell state method.
Tensor_t fDy
cached activation gradient (input of backward) as T x B x S
Matrix_t & GetCandidateDerivativesAt(size_t i)
std::vector< Matrix_t > fDerivativesUpdate
First fDerivatives of the activations update gate.
size_t GetTimeSteps() const
const Matrix_t & GetWeightsUpdateGateState() const
std::vector< Matrix_t > & GetDerivativesCandidate()
bool DoesRememberState() const
const Matrix_t & GetWeightsResetGateState() const
void CandidateValue(const Matrix_t &input, Matrix_t &dc)
Decides the new candidate values (NN with Tanh)
TBasicGRULayer(size_t batchSize, size_t stateSize, size_t inputSize, size_t timeSteps, bool rememberState=false, bool returnSequence=false, bool resetGateAfter=false, DNN::EActivationFunction f1=DNN::EActivationFunction::kSigmoid, DNN::EActivationFunction f2=DNN::EActivationFunction::kTanh, bool training=true, DNN::EInitialization fA=DNN::EInitialization::kZero)
Constructor.
Matrix_t & GetUpdateDerivativesAt(size_t i)
Matrix_t & fWeightsUpdateGate
Update Gate weights for input, fWeights[2].
typename Architecture_t::DropoutDescriptor_t HelperDescriptor_t
Matrix_t & GetResetGateValue()
Generic General Layer class.
virtual void Initialize()
Initialize the weights and biases according to the given initialization method.
size_t GetInputWidth() const
XMLNodePointer_t NewChild(XMLNodePointer_t parent, XMLNsPointer_t ns, const char *name, const char *content=nullptr)
create new child element for parent node
XMLAttrPointer_t NewAttr(XMLNodePointer_t xmlnode, XMLNsPointer_t, const char *name, const char *value)
creates new attribute for xmlnode, namespaces are not supported for attributes
EActivationFunction
Enum that represents layer activation functions.
create variable transformations