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;
142 TBasicGRULayer(
size_t batchSize,
size_t stateSize,
size_t inputSize,
143 size_t timeSteps,
bool rememberState =
false,
bool returnSequence =
false,
144 bool resetGateAfter =
false,
169 const Tensor_t &activations_backward)
override;
177 const Matrix_t & precStateActivations,
193 void Print()
const override;
309template <
typename Architecture_t>
314 :
VGeneralLayer<Architecture_t>(batchSize, 1, timeSteps, inputSize, 1, (returnSequence) ? timeSteps : 1, stateSize,
315 6, {stateSize, stateSize, stateSize, stateSize, stateSize, stateSize},
316 {inputSize, inputSize, inputSize, stateSize, stateSize, stateSize}, 3,
317 {stateSize, stateSize, stateSize}, {1, 1, 1}, batchSize,
318 (returnSequence) ? timeSteps : 1, stateSize, fA),
319 fStateSize(stateSize), fTimeSteps(timeSteps), fRememberState(rememberState), fReturnSequence(returnSequence), fResetGateAfter(resetGateAfter),
320 fF1(
f1), fF2(f2), fResetValue(batchSize, stateSize), fUpdateValue(batchSize, stateSize),
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))
333 for (
size_t i = 0; i < timeSteps; ++i) {
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>
432 Matrix_t tmpState(fResetValue.GetNrows(), fResetValue.GetNcols());
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>
450 Matrix_t tmpState(fUpdateValue.GetNrows(), fUpdateValue.GetNcols());
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()) {
503 assert(
input.GetStrides()[1] == this->GetInputSize());
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;
521 Architecture_t::RNNForward(
x, hx, cx, weights,
y, hy, cy, rnnDesc, rnnWork, isTraining);
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());
543 Architecture_t::Rearrange(arrInput,
input);
545 Tensor_t arrOutput ( fTimeSteps, this->GetBatchSize(), fStateSize );
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);
571 Matrix_t arrOutputMt = arrOutput[t];
572 Architecture_t::Copy(arrOutputMt, fState);
576 Architecture_t::Rearrange(this->GetOutput(), arrOutput);
579 Tensor_t tmp = arrOutput.At(fTimeSteps - 1);
582 tmp = tmp.Reshape({tmp.GetShape()[0], tmp.GetShape()[1], 1});
583 assert(tmp.GetSize() == this->GetOutput().GetSize());
584 assert(tmp.GetShape()[0] == this->GetOutput().GetShape()[2]);
585 Architecture_t::Rearrange(this->GetOutput(), tmp);
592template <
typename Architecture_t>
596 Architecture_t::Hadamard(fState, updateGateValues);
600 for (
size_t j = 0; j < (size_t) tmp.GetNcols(); j++) {
601 for (
size_t i = 0; i < (size_t) tmp.GetNrows(); i++) {
602 tmp(i,j) = 1 - tmp(i,j);
607 Architecture_t::Hadamard(candidateValues, tmp);
608 Architecture_t::ScaleAdd(fState, candidateValues);
612template <
typename Architecture_t>
614 const Tensor_t &activations_backward)
618 if (Architecture_t::IsCudnn()) {
626 assert(activations_backward.GetStrides()[1] == this->GetInputSize());
629 Architecture_t::Rearrange(
x, activations_backward);
631 if (!fReturnSequence) {
634 Architecture_t::InitializeZero(dy);
637 Tensor_t tmp2 = dy.At(dy.GetShape()[0] - 1).Reshape({dy.GetShape()[1], 1, dy.GetShape()[2]});
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();
650 auto &weightGradients = this->GetWeightGradientsTensor();
654 Architecture_t::InitializeZero(weightGradients);
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);
672 if (gradients_backward.GetSize() != 0)
673 Architecture_t::Rearrange(gradients_backward, dx);
681 Matrix_t state_gradients_backward(this->GetBatchSize(), fStateSize);
686 if (gradients_backward.GetSize() == 0 || gradients_backward[0].GetNrows() == 0 || gradients_backward[0].GetNcols() == 0) {
690 Tensor_t arr_gradients_backward ( fTimeSteps, this->GetBatchSize(), this->GetInputSize());
695 Tensor_t arr_activations_backward ( fTimeSteps, this->GetBatchSize(), this->GetInputSize());
697 Architecture_t::Rearrange(arr_activations_backward, activations_backward);
701 Tensor_t arr_output ( fTimeSteps, this->GetBatchSize(), fStateSize);
703 Matrix_t initState(this->GetBatchSize(), fStateSize);
707 Tensor_t arr_actgradients ( fTimeSteps, this->GetBatchSize(), fStateSize);
709 if (fReturnSequence) {
710 Architecture_t::Rearrange(arr_output, this->GetOutput());
711 Architecture_t::Rearrange(arr_actgradients, this->GetActivationGradients());
715 Architecture_t::InitializeZero(arr_actgradients);
717 Tensor_t tmp_grad = arr_actgradients.At(fTimeSteps - 1).Reshape({this->GetBatchSize(), fStateSize, 1});
718 assert(tmp_grad.GetSize() == this->GetActivationGradients().GetSize());
719 assert(tmp_grad.GetShape()[0] ==
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--) {
746 Architecture_t::ScaleAdd(state_gradients_backward, arr_actgradients[t-1]);
748 const Matrix_t &prevStateActivations = arr_output[t-2];
749 Matrix_t dx = arr_gradients_backward[t-1];
751 CellBackward(state_gradients_backward, prevStateActivations,
752 this->GetResetGateTensorAt(t-1), this->GetUpdateGateTensorAt(t-1),
753 this->GetCandidateGateTensorAt(t-1),
754 arr_activations_backward[t-1], dx ,
755 fDerivativesReset[t-1], fDerivativesUpdate[t-1],
756 fDerivativesCandidate[t-1]);
758 const Matrix_t &prevStateActivations = initState;
759 Matrix_t dx = arr_gradients_backward[t-1];
760 CellBackward(state_gradients_backward, prevStateActivations,
761 this->GetResetGateTensorAt(t-1), this->GetUpdateGateTensorAt(t-1),
762 this->GetCandidateGateTensorAt(t-1),
763 arr_activations_backward[t-1], dx ,
764 fDerivativesReset[t-1], fDerivativesUpdate[t-1],
765 fDerivativesCandidate[t-1]);
770 Architecture_t::Rearrange(gradients_backward, arr_gradients_backward );
777template <
typename Architecture_t>
779 const Matrix_t & precStateActivations,
788 return Architecture_t::GRULayerBackward(state_gradients_backward,
789 fWeightsResetGradients, fWeightsUpdateGradients, fWeightsCandidateGradients,
790 fWeightsResetStateGradients, fWeightsUpdateStateGradients,
791 fWeightsCandidateStateGradients, fResetBiasGradients, fUpdateBiasGradients,
792 fCandidateBiasGradients, dr, du, dc,
793 precStateActivations,
794 reset_gate, update_gate, candidate_gate,
795 fWeightsResetGate, fWeightsUpdateGate, fWeightsCandidate,
796 fWeightsResetGateState, fWeightsUpdateGateState, fWeightsCandidateState,
797 input, input_gradient, fResetGateAfter);
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));
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 GetBatchSize() const
Getters.
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