From d45a955a3c828f7698918ac0caedaebf71dde1b2 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Sun, 2 Feb 2025 06:10:53 -0500 Subject: [PATCH] Allow ragged sequences for RNN training and prediction (#3870) * Add LinearRecurrent layer and tests. * Update RecurrentLayer for correct BPTT and document API changes. * Reimplement BPTT---correctly, this time. * Fix RNN implementation and add tests. * Fix and test LSTM. * Fix BPTT implementation to use all time steps and add some more tests. * Fix and tune recurrent sine tests. * Fix serialization for LinearRecurrent. * Fix compilation warning when serialization is not enabled. * When I use the MSE I expect the mean squared error, not the squared error. These reductions will need to be revisited at some point. * Restore Reber grammar tests, revamp them, and tune other tests. * Fix a few tests. * Update HISTORY. * Fix style issues. * Fix one small behavioral change. * Small patch to support ragged sequence lengths. * Clean up ragged sequence code and add tests. * Revert inadvertent changes to NegativeLogLikelihood. * Fix style issues. * Fix conditions so that we can serialize NNs even if we didn't define MLPACK_ENABLE_ANN_SERIALIZATION. --- src/mlpack/methods/ann/ffn_impl.hpp | 11 +- src/mlpack/methods/ann/rnn.hpp | 104 ++++++++++++- src/mlpack/methods/ann/rnn_impl.hpp | 140 ++++++++++++++++-- src/mlpack/tests/ann/layer/lstm.cpp | 8 +- .../tests/ann/recurrent_network_test.cpp | 75 ++++++++++ 5 files changed, 309 insertions(+), 29 deletions(-) diff --git a/src/mlpack/methods/ann/ffn_impl.hpp b/src/mlpack/methods/ann/ffn_impl.hpp index 4465a69bed..bcaf24c360 100644 --- a/src/mlpack/methods/ann/ffn_impl.hpp +++ b/src/mlpack/methods/ann/ffn_impl.hpp @@ -375,16 +375,15 @@ void FFN< MatType >::serialize(Archive& ar, const uint32_t /* version */) { - #ifndef MLPACK_ENABLE_ANN_SERIALIZATION + #if !defined(MLPACK_ENABLE_ANN_SERIALIZATION) && \ + !defined(MLPACK_ANN_IGNORE_SERIALIZATION_WARNING) // Note: if you define MLPACK_IGNORE_ANN_SERIALIZATION_WARNING, you had // better ensure that every layer you are serializing has had // CEREAL_REGISTER_TYPE() called somewhere. See layer/serialization.hpp for // more information. - #ifndef MLPACK_ANN_IGNORE_SERIALIZATION_WARNING - throw std::runtime_error("Cannot serialize a neural network unless " - "MLPACK_ENABLE_ANN_SERIALIZATION is defined! See the \"Additional " - "build options\" section of the README for more information."); - #endif + throw std::runtime_error("Cannot serialize a neural network unless " + "MLPACK_ENABLE_ANN_SERIALIZATION is defined! See the \"Additional " + "build options\" section of the README for more information."); (void) ar; #else diff --git a/src/mlpack/methods/ann/rnn.hpp b/src/mlpack/methods/ann/rnn.hpp index 1db25d1b90..ee97d6f742 100644 --- a/src/mlpack/methods/ann/rnn.hpp +++ b/src/mlpack/methods/ann/rnn.hpp @@ -160,6 +160,73 @@ class RNN arma::Cube responses, CallbackTypes&&... callbacks); + /** + * Train the recurrent network on the given input data using the given + * optimizer, given that input sequences may have different lengths. + * + * This will use the existing model parameters as a starting point for the + * optimization. If this is not what you want, then you should access the + * parameters vector directly with Parameters() and modify it as desired. + * + * Note that due to shuffling, training will make a copy of the data, unless + * you use `std::move()` to pass the `predictors` and `responses` (that is, + * `Train(std::move(predictors), std::move(responses))`). + * + * @tparam OptimizerType Type of optimizer to use to train the model. + * @tparam CallbackTypes Types of Callback Functions. + * @param predictors Input training variables. + * @param responses Outputs results from input training variables. + * @param sequenceLengths Length of each input sequences. Should have size + * `predictors.n_cols`, and all values should be less than or equal to + * `predictors.n_slices`. + * @param optimizer Instantiated optimizer used to train the model. + * @param callbacks Callback function for ensmallen optimizer `OptimizerType`. + * See https://www.ensmallen.org/docs.html#callback-documentation. + * @return The final objective of the trained model (NaN or Inf on error). + */ + template + typename MatType::elem_type Train( + arma::Cube predictors, + arma::Cube responses, + arma::urowvec sequenceLengths, + OptimizerType& optimizer, + CallbackTypes&&... callbacks); + + /** + * Train the recurrent network on the given input data, given that each input + * sequence may have a different length. By default, the RMSProp optimization + * algorithm is used, but others can be specified (such as ens::SGD). + * + * When passing sequences with different lengths, the batch size of the + * optimizer must be set to 1; if it is not, an exception will be thrown + * during training. + * + * This will use the existing model parameters as a starting point for the + * optimization. If this is not what you want, then you should access the + * parameters vector directly with Parameters() and modify it as desired. + * + * Note that due to shuffling, training will make a copy of the data, unless + * you use `std::move()` to pass the `predictors` and `responses` (that is, + * `Train(std::move(predictors), std::move(responses))`). + * + * @tparam OptimizerType Type of optimizer to use to train the model. + * @tparam CallbackTypes Types of Callback Functions. + * @param predictors Input training variables. + * @param responses Outputs results from input training variables. + * @param sequenceLengths Length of each input sequences. Should have size + * `predictors.n_cols`, and all values should be less than or equal to + * `predictors.n_slices`. + * @param callbacks Callback function for ensmallen optimizer `OptimizerType`. + * See https://www.ensmallen.org/docs.html#callback-documentation. + * @return The final objective of the trained model (NaN or Inf on error). + */ + template + typename MatType::elem_type Train( + arma::Cube predictors, + arma::Cube responses, + arma::urowvec sequenceLengths, + CallbackTypes&&... callbacks); + /** * Predict the responses to a given set of predictors. The responses will * reflect the output of the given output layer as returned by the @@ -173,6 +240,24 @@ class RNN arma::Cube& results, const size_t batchSize = 128); + /** + * Predict the responses to a given set of predictors, given that each + * sequence can have a different length. The responses will reflect the output + * of the given output layer as returned by the output layer function. + * + * Slices of column `i` of `results` at time indexes greater than + * `sequenceLengths[i]` should not be considered valid predictions. + * + * The batch size is limited to 1 when predicting on sequences of different + * lengths. + * + * @param predictors Input predictors. + * @param results Matrix to put output predictions of responses into. + */ + void Predict(const arma::Cube& predictors, + arma::Cube& results, + const arma::urowvec& sequenceLengths); + // Return the nujmber of weights in the model. size_t WeightSize() { return network.WeightSize(); } @@ -337,9 +422,12 @@ class RNN * * @param predictors Input data variables. * @param responses Outputs results from input data variables. + * @param sequenceLengths (Optional) sequence length for each predictor + * sequence. */ void ResetData(arma::Cube predictors, - arma::Cube responses); + arma::Cube responses, + arma::urowvec sequenceLengths = arma::urowvec()); private: // Helper functions. @@ -365,14 +453,18 @@ class RNN //! occasionally resetting any memory cells. FFN network; - //! The matrix of data points (predictors). This member is empty, except - //! during training---we must store a local copy of the training data since - //! the ensmallen optimizer will not provide training data. + // The matrix of data points (predictors). These members are empty, except + // during training---we must store a local copy of the training data since + // the ensmallen optimizer will not provide training data. arma::Cube predictors; - //! The matrix of responses to the input data points. This member is empty, - //! except during training. + // The matrix of responses to the input data points. This member is empty, + // except during training. arma::Cube responses; + + // The length of each input sequence. If this is empty, then every sequence + // is assuemd to have the same length (`predictors.n_slices`). + arma::urowvec sequenceLengths; }; // class RNNType } // namespace mlpack diff --git a/src/mlpack/methods/ann/rnn_impl.hpp b/src/mlpack/methods/ann/rnn_impl.hpp index 0d0e44b6bd..6ebb32f4ab 100644 --- a/src/mlpack/methods/ann/rnn_impl.hpp +++ b/src/mlpack/methods/ann/rnn_impl.hpp @@ -161,7 +161,7 @@ typename MatType::elem_type RNN< OptimizerType& optimizer, CallbackTypes&&... callbacks) { - ResetData(std::move(predictors), std::move(responses)); + ResetData(std::move(predictors), std::move(responses), arma::urowvec()); network.WarnMessageMaxIterations(optimizer, this->predictors.n_cols); @@ -199,6 +199,63 @@ typename MatType::elem_type RNN< callbacks...); } +template< + typename OutputLayerType, + typename InitializationRuleType, + typename MatType +> +template +typename MatType::elem_type RNN< + OutputLayerType, + InitializationRuleType, + MatType +>::Train( + arma::Cube predictors, + arma::Cube responses, + arma::urowvec sequenceLengths, + OptimizerType& optimizer, + CallbackTypes&&... callbacks) +{ + ResetData(std::move(predictors), std::move(responses), + std::move(sequenceLengths)); + + network.WarnMessageMaxIterations(optimizer, this->predictors.n_cols); + + // Ensure that the network can be used. + network.CheckNetwork("RNN::Train()", this->predictors.n_rows, true, true); + + // Train the model. + Timer::Start("rnn_optimization"); + const typename MatType::elem_type out = + optimizer.Optimize(*this, network.Parameters(), callbacks...); + Timer::Stop("rnn_optimization"); + + Log::Info << "RNN::Train(): final objective of trained model is " << out + << "." << std::endl; + return out; +} + +template< + typename OutputLayerType, + typename InitializationRuleType, + typename MatType +> +template +typename MatType::elem_type RNN< + OutputLayerType, + InitializationRuleType, + MatType +>::Train( + arma::Cube predictors, + arma::Cube responses, + arma::urowvec sequenceLengths, + CallbackTypes&&... callbacks) +{ + OptimizerType optimizer; + return Train(std::move(predictors), std::move(responses), + std::move(sequenceLengths), optimizer, callbacks...); +} + template< typename OutputLayerType, typename InitializationRuleType, @@ -246,6 +303,51 @@ void RNN< } } +template< + typename OutputLayerType, + typename InitializationRuleType, + typename MatType +> +void RNN< + OutputLayerType, + InitializationRuleType, + MatType +>::Predict( + const arma::Cube& predictors, + arma::Cube& results, + const arma::urowvec& sequenceLengths) +{ + // Ensure that the network is configured correctly. + network.CheckNetwork("RNN::Predict()", predictors.n_rows, true, false); + + results.set_size(network.network.OutputSize(), predictors.n_cols, + single ? 1 : predictors.n_slices); + + MatType inputAlias, outputAlias; + for (size_t i = 0; i < predictors.n_cols; i++) + { + // Since we aren't doing a backward pass, we don't actually need to store + // the state for each time step---we can fit it all in one buffer. + ResetMemoryState(0, 1); + + // Iterate over all time steps. + const size_t steps = sequenceLengths[i]; + for (size_t t = 0; t < steps; ++t) + { + SetCurrentStep(t, (t == steps - 1)); + + // Create aliases for the input and output. If we are in single mode, we + // always output into the same slice. + MakeAlias(inputAlias, predictors.slice(t), predictors.n_rows, 1, + i * predictors.n_rows); + MakeAlias(outputAlias, results.slice(single ? 0 : t), results.n_rows, 1, + i * results.n_rows); + + network.Forward(inputAlias, outputAlias); + } + } +} + template< typename OutputLayerType, typename InitializationRuleType, @@ -285,16 +387,15 @@ void RNN< MatType >::serialize(Archive& ar, const uint32_t /* version */) { - #ifndef MLPACK_ENABLE_ANN_SERIALIZATION + #if !defined(MLPACK_ENABLE_ANN_SERIALIZATION) && \ + !defined(MLPACK_ANN_IGNORE_SERIALIZATION_WARNING) // Note: if you define MLPACK_IGNORE_ANN_SERIALIZATION_WARNING, you had // better ensure that every layer you are serializing has had // CEREAL_REGISTER_TYPE() called somewhere. See layer/serialization.hpp for // more information. - #ifndef MLPACK_IGNORE_ANN_SERIALIZATION_WARNING - throw std::runtime_error("Cannot serialize a neural network unless " - "MLPACK_ENABLE_ANN_SERIALIZATION is defined! See the \"Additional " - "build options\" section of the README for more information."); - #endif + throw std::runtime_error("Cannot serialize a neural network unless " + "MLPACK_ENABLE_ANN_SERIALIZATION is defined! See the \"Additional " + "build options\" section of the README for more information."); (void) ar; #else @@ -308,6 +409,7 @@ void RNN< // middle of training and resume. predictors.clear(); responses.clear(); + sequenceLengths.clear(); } #endif } @@ -335,13 +437,18 @@ typename MatType::elem_type RNN< ResetMemoryState(1, batchSize); MatType output(network.network.OutputSize(), batchSize); + if (sequenceLengths.n_elem > 0 && batchSize != 1) + throw std::invalid_argument("Batch size must be 1 for ragged sequences!"); + typename MatType::elem_type loss = 0.0; MatType stepData, responseData; - for (size_t t = 0; t < predictors.n_slices; ++t) + const size_t steps = (sequenceLengths.n_elem == 0) ? predictors.n_slices : + sequenceLengths[begin]; + for (size_t t = 0; t < steps; ++t) { // Manually reset the data of the network to be an alias of the current time // step. - SetCurrentStep(t, (t == predictors.n_slices - 1)); + SetCurrentStep(t, (t == steps)); MakeAlias(network.predictors, predictors.slice(t), predictors.n_rows, batchSize, begin * predictors.slice(t).n_rows); const size_t responseStep = (single) ? 0 : t; @@ -390,6 +497,9 @@ typename MatType::elem_type RNN< { network.CheckNetwork("RNN::EvaluateWithGradient()", predictors.n_rows); + if (sequenceLengths.n_elem > 0 && batchSize != 1) + throw std::invalid_argument("Batch size must be 1 for ragged sequences!"); + typename MatType::elem_type loss = 0; // We must save anywhere between 1 and `bpttSteps` states, but we are limited @@ -418,9 +528,11 @@ typename MatType::elem_type RNN< // For backpropagation through time, we must backpropagate for every // subsequence of length `bpttSteps`. Before we've taken `bpttSteps` though, // we will be backpropagating shorter sequences. - for (size_t t = 0; t < predictors.n_slices; ++t) + const size_t steps = (sequenceLengths.n_elem == 0) ? predictors.n_slices : + sequenceLengths[begin]; + for (size_t t = 0; t < steps; ++t) { - SetCurrentStep(t, (t == (predictors.n_slices - 1))); + SetCurrentStep(t, (t == (steps - 1))); // Make an alias of the step's data for the forward pass. MakeAlias(stepData, predictors.slice(t), predictors.n_rows, batchSize, @@ -431,7 +543,7 @@ typename MatType::elem_type RNN< // Determine what the response should be. If we are in single mode but not // at the end of the sequence, we don't do a backwards pass. - if (single && t != responses.n_slices - 1) + if (single && t != steps - 1) { continue; } @@ -535,10 +647,12 @@ void RNN< MatType >::ResetData( arma::Cube predictors, - arma::Cube responses) + arma::Cube responses, + arma::urowvec sequenceLengths) { this->predictors = std::move(predictors); this->responses = std::move(responses); + this->sequenceLengths = std::move(sequenceLengths); } template< diff --git a/src/mlpack/tests/ann/layer/lstm.cpp b/src/mlpack/tests/ann/layer/lstm.cpp index 28a459b8df..4b8d30740b 100644 --- a/src/mlpack/tests/ann/layer/lstm.cpp +++ b/src/mlpack/tests/ann/layer/lstm.cpp @@ -93,10 +93,10 @@ void CheckRNNCopyFunction(ModelType* network1, network1->Predict(trainData, predictions1); RNN<> network2 = *network1; - // Deallocate all of network1's memory, so we can check that network2 does not - // use any of that memory. delete network1; + // Deallocating all of network1's memory, so that network2 does not use any + // of that memory. network2.Predict(trainData, predictions2); CheckMatrices(predictions1, predictions2); } @@ -116,10 +116,10 @@ void CheckRNNMoveFunction(ModelType* network1, network1->Predict(trainData, predictions1); RNN<> network2(std::move(*network1)); - // Deallocate all of network1's memory, so we can check that network2 does not - // use any of that memory. delete network1; + // Deallocating all of network1's memory, so that network2 does not use any + // of that memory. network2.Predict(trainData, predictions2); CheckMatrices(predictions1, predictions2); } diff --git a/src/mlpack/tests/ann/recurrent_network_test.cpp b/src/mlpack/tests/ann/recurrent_network_test.cpp index 029de4235e..431a446085 100644 --- a/src/mlpack/tests/ann/recurrent_network_test.cpp +++ b/src/mlpack/tests/ann/recurrent_network_test.cpp @@ -1031,3 +1031,78 @@ TEST_CASE("LSTMEmbeddedReberGrammarTest", "[RecurrentNetworkTest]") model.Add(); ReberGrammarTestNetwork(model, true); } + +/** + * Test that we can train an RNN on sequences of different lengths, and get + * roughly the same thing we would for training on non-ragged sequences. + */ +TEST_CASE("RNNRaggedSequenceTest", "[RecurrentNetworkTest]") +{ + const size_t rho = 25; + const size_t numEpochs = 3; + + // Generate noisy sine data. + arma::cube data, responses; + GenerateNoisySinRNN(data, responses, 500, rho + 35); + arma::cube origData = data; + arma::cube origResponses = responses; + + // Assign random sequence lengths for each sine. + arma::urowvec lengths = arma::randi(500, distr_param(40, 60)); + + // Set garbage data for anything past the end of a sequence. + for (size_t c = 0; c < 500; ++c) + { + if (lengths[c] == 60) + continue; + + data.subcube(0, c, lengths[c], + data.n_rows - 1, c, data.n_slices - 1).randu(); + responses.subcube(0, c, lengths[c], + responses.n_rows - 1, c, responses.n_slices - 1).randu(); + } + + // Build a network and train it. + RMSProp opt(0.003, 1, 0.99, 1e-08, 500 * numEpochs, 1e-5); + + RNN net(rho); + net.Add(10); + net.Add(1); + + // Train on all the data. + net.Train(data, responses, lengths, opt); + + // Make sure that the predictions match the data reasonably. + arma::cube prediction; + net.Predict(data, prediction, lengths); + + // Sum the error for all sequences. + size_t timeSteps = 0; + double totalError = 0.0; + for (size_t c = 0; c < 500; ++c) + { + timeSteps += lengths[c]; + totalError += accu(abs(vectorise(responses.subcube( + 0, c, 0, responses.n_rows - 1, c, lengths[c] - 1)) - + vectorise(prediction.subcube( + 0, c, 0, prediction.n_rows - 1, c, lengths[c] - 1)))); + } + + const double averageError = (totalError / timeSteps); + + // Now compute another network where we don't use the sequence lengths. + RNN net2(rho); + net2.Add(10); + net2.Add(1); + + // Train and predict, then compute the sum error. + RMSProp opt2(0.003, 1, 0.99, 1e-08, 500 * numEpochs / 2, 1e-5); + net2.Train(origData, origResponses, opt2); + net2.Predict(origData, prediction); + const double refAverageError = mean(abs(vectorise(origResponses) - + vectorise(prediction))); + + // There can be some margin in the results because we are not training on as + // much data for the ragged sequences. + REQUIRE(abs(averageError - refAverageError) <= 0.1); +}