Merge pull request #1800 from MuLx10/LSTMCell

Added cellState as output params in LSTM.
This commit is contained in:
Shikhar Jaiswal
2019-04-06 13:04:57 +05:30
committed by GitHub
3 changed files with 219 additions and 7 deletions
+17 -7
View File
@@ -19,12 +19,7 @@ namespace mlpack {
namespace ann /** Artificial Neural Network. */ {
/**
* An implementation of a lstm network layer.
*
* This class allows specification of the type of the activation functions used
* for the gates and cells and also of the type of the function used to
* initialize and update the peephole weights.
* Implementation of the LSTM module class.
* The implementation corresponds to the following algorithm:
*
* @f{eqnarray}{
@@ -78,7 +73,7 @@ class LSTM
const size_t rho = std::numeric_limits<size_t>::max());
/**
* Ordinary feed forward pass of a neural network, evaluating the function
* Ordinary feed-forward pass of a neural network, evaluating the function
* f(x) by propagating the activity forward through f.
*
* @param input Input data used for evaluating the specified function.
@@ -87,6 +82,21 @@ class LSTM
template<typename InputType, typename OutputType>
void Forward(InputType&& input, OutputType&& output);
/**
* Ordinary feed-forward pass of a neural network, evaluating the function
* f(x) by propagating the activity forward through f.
*
* @param input Input data used for evaluating the specified function.
* @param output Resulting output activation.
* @param cellState Cell state of the LSTM.
* @param useCellState Use the cellState passed in the LSTM cell.
*/
template<typename InputType, typename OutputType>
void Forward(InputType&& input,
OutputType&& output,
OutputType&& cellState,
bool useCellState = false);
/**
* Ordinary feed backward pass of a neural network, calculating the function
* f(x) by propagating x backwards trough f. Using the results from the feed
@@ -159,10 +159,24 @@ void LSTM<InputDataType, OutputDataType>::Reset()
offset, outSize, 1, false, false);
}
// Forward when cellState is not needed.
template<typename InputDataType, typename OutputDataType>
template<typename InputType, typename OutputType>
void LSTM<InputDataType, OutputDataType>::Forward(
InputType&& input, OutputType&& output)
{
//! Locally-stored cellState.
OutputType cellState;
Forward(std::move(input), std::move(output), std::move(cellState), false);
}
// Forward when cellState is needed overloaded LSTM::Forward().
template<typename InputDataType, typename OutputDataType>
template<typename InputType, typename OutputType>
void LSTM<InputDataType, OutputDataType>::Forward(InputType&& input,
OutputType&& output,
OutputType&& cellState,
bool useCellState)
{
// Check if the batch size changed, the number of cols is defines the input
// batch size.
@@ -187,6 +201,18 @@ void LSTM<InputDataType, OutputDataType>::Forward(
if (forwardStep > 0)
{
if (useCellState)
{
if (!cellState.is_empty())
{
cell.cols(forwardStep - batchSize,
forwardStep - batchSize + batchStep) = cellState;
}
else
{
throw std::runtime_error("Cell parameter is empty.");
}
}
inputGate.cols(forwardStep, forwardStep + batchStep) +=
arma::repmat(cell2GateInputWeight, 1, batchSize) %
cell.cols(forwardStep - batchSize, forwardStep - batchSize + batchStep);
@@ -249,6 +275,9 @@ void LSTM<InputDataType, OutputDataType>::Forward(
output = OutputType(outParameter.memptr() +
(forwardStep + batchSize) * outSize, outSize, batchSize, false, false);
cellState = OutputType(cell.memptr() +
forwardStep * outSize, outSize, batchSize, false, false);
forwardStep += batchSize;
if ((forwardStep / batchSize) == bpttSteps)
{
+173
View File
@@ -901,6 +901,179 @@ BOOST_AUTO_TEST_CASE(GradientFastLSTMLayerTest)
BOOST_REQUIRE_LE(CheckGradient(function), 0.2);
}
/**
* Testing the overloaded Forward() of the LSTM layer, for retrieving the cell
* state. Besides output, the overloaded function provides read access to cell
* state of the LSTM layer.
*/
BOOST_AUTO_TEST_CASE(ReadCellStateParamLSTMLayerTest)
{
const size_t rho = 5, inputSize = 3, outputSize = 2;
// Provide input of all ones.
arma::cube input = arma::ones(inputSize, outputSize, rho);
arma::mat inputGate, forgetGate, outputGate, hidden;
arma::mat outLstm, cellLstm;
// LSTM layer.
LSTM<> lstm(inputSize, outputSize, rho);
lstm.Reset();
lstm.ResetCell(rho);
// Initialize the weights to all ones.
lstm.Parameters().ones();
arma::mat inputWeight = arma::ones(outputSize, inputSize);
arma::mat outputWeight = arma::ones(outputSize, outputSize);
arma::mat bias = arma::ones(outputSize, input.n_cols);
arma::mat cellCalc = arma::zeros(outputSize, input.n_cols);
arma::mat outCalc = arma::zeros(outputSize, input.n_cols);
for (size_t seqNum = 0; seqNum < rho; ++seqNum)
{
// Wrap a matrix around our data to avoid a copy.
arma::mat stepData(input.slice(seqNum).memptr(),
input.n_rows, input.n_cols, false, true);
// Apply Forward() on LSTM layer.
lstm.Forward(std::move(stepData), // Input.
std::move(outLstm), // Output.
std::move(cellLstm), // Cell state.
false); // Don't write into the cell state.
// Compute the value of cell state and output.
// i = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
inputGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// f = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
forgetGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// z = tanh(W.dot(x) + W.dot(h) + b).
hidden = arma::tanh(inputWeight * stepData +
outputWeight * outCalc + bias);
// c = f * c + i * z.
cellCalc = forgetGate % cellCalc + inputGate % hidden;
// o = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
outputGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// h = o * tanh(c).
outCalc = outputGate % arma::tanh(cellCalc);
CheckMatrices(outLstm, outCalc, 1e-12);
CheckMatrices(cellLstm, cellCalc, 1e-12);
}
}
/**
* Testing the overloaded Forward() of the LSTM layer, for retrieving the cell
* state. Besides output, the overloaded function provides write access to cell
* state of the LSTM layer.
*/
BOOST_AUTO_TEST_CASE(WriteCellStateParamLSTMLayerTest)
{
const size_t rho = 5, inputSize = 3, outputSize = 2;
// Provide input of all ones.
arma::cube input = arma::ones(inputSize, outputSize, rho);
arma::mat inputGate, forgetGate, outputGate, hidden;
arma::mat outLstm, cellLstm;
arma::mat cellCalc;
// LSTM layer.
LSTM<> lstm(inputSize, outputSize, rho);
lstm.Reset();
lstm.ResetCell(rho);
// Initialize the weights to all ones.
lstm.Parameters().ones();
arma::mat inputWeight = arma::ones(outputSize, inputSize);
arma::mat outputWeight = arma::ones(outputSize, outputSize);
arma::mat bias = arma::ones(outputSize, input.n_cols);
arma::mat outCalc = arma::zeros(outputSize, input.n_cols);
for (size_t seqNum = 0; seqNum < rho; ++seqNum)
{
// Wrap a matrix around our data to avoid a copy.
arma::mat stepData(input.slice(seqNum).memptr(),
input.n_rows, input.n_cols, false, true);
if (cellLstm.is_empty())
{
// Set the cell state to zeros.
cellLstm = arma::zeros(outputSize, input.n_cols);
cellCalc = arma::zeros(outputSize, input.n_cols);
}
else
{
// Set the cell state to zeros.
cellLstm = arma::zeros(cellLstm.n_rows, cellLstm.n_cols);
cellCalc = arma::zeros(cellCalc.n_rows, cellCalc.n_cols);
}
// Apply Forward() on the LSTM layer.
lstm.Forward(std::move(stepData), // Input.
std::move(outLstm), // Output.
std::move(cellLstm), // Cell state.
true); // Write into cell state.
// Compute the value of cell state and output.
// i = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
inputGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// f = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
forgetGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// z = tanh(W.dot(x) + W.dot(h) + b).
hidden = arma::tanh(inputWeight * stepData +
outputWeight * outCalc + bias);
// c = f * c + i * z.
cellCalc = forgetGate % cellCalc + inputGate % hidden;
// o = sigmoid(W.dot(x) + W.dot(h) + W.dot(c) + b).
outputGate = 1.0 /(1 + arma::exp(-(inputWeight * stepData +
outputWeight * outCalc + outputWeight % cellCalc + bias)));
// h = o * tanh(c).
outCalc = outputGate % arma::tanh(cellCalc);
CheckMatrices(outLstm, outCalc, 1e-12);
CheckMatrices(cellLstm, cellCalc, 1e-12);
}
// Attempting to write empty matrix into cell state.
lstm.Reset();
lstm.ResetCell(rho);
arma::mat stepData(input.slice(0).memptr(),
input.n_rows, input.n_cols, false, true);
lstm.Forward(std::move(stepData), // Input.
std::move(outLstm), // Output.
std::move(cellLstm), // Cell state.
true); // Write into cell state.
for (size_t seqNum = 1; seqNum < rho; ++seqNum)
{
arma::mat empty;
// Should throw error.
BOOST_REQUIRE_THROW(lstm.Forward(std::move(stepData), // Input.
std::move(outLstm), // Output.
std::move(empty), // Cell state.
true), // Write into cell state.
std::runtime_error);
}
}
/**
* Check if the gradients computed by GRU cell are close enough to the
* approximation of the gradients.