diff --git a/src/mlpack/methods/ann/augmented/tasks/add_impl.hpp b/src/mlpack/methods/ann/augmented/tasks/add_impl.hpp index 13cec239ee..09f35f84e1 100644 --- a/src/mlpack/methods/ann/augmented/tasks/add_impl.hpp +++ b/src/mlpack/methods/ann/augmented/tasks/add_impl.hpp @@ -116,7 +116,8 @@ void AddTask::Generate(arma::mat& input, arma::mat& labels, } } -void Binarize(const arma::field& input, arma::field& output) +void AddTask::Binarize(const arma::field& input, + arma::field& output) { arma::field procData(input.n_elem); for (size_t i = 0; i < input.n_elem; ++i) diff --git a/src/mlpack/tests/augmented_rnns_tasks_test.cpp b/src/mlpack/tests/augmented_rnns_tasks_test.cpp index da7885b733..fb126ec10b 100644 --- a/src/mlpack/tests/augmented_rnns_tasks_test.cpp +++ b/src/mlpack/tests/augmented_rnns_tasks_test.cpp @@ -89,17 +89,18 @@ class HardCodedCopyModel { class HardCodedSortModel { public: - HardCodedSortModel() {} + HardCodedSortModel(size_t bitLen) : bitLen(bitLen) {} void Train(arma::field& predictors, arma::field& labels) { assert(predictors.n_elem == labels.n_elem); - bitLen = predictors.at(0).n_rows; } void Predict( arma::mat& predictors, arma::mat& labels) { + predictors = predictors.t(); + predictors.reshape(bitLen, predictors.n_elem / bitLen); size_t len = predictors.n_cols; labels.zeros(bitLen, len); vector> vals(len); @@ -115,6 +116,7 @@ class HardCodedSortModel { for (size_t j = 0; j < len; ++j) { labels.col(j) = predictors.col(vals[j].second); } + labels.reshape(predictors.n_elem, 1); } void Predict(arma::field& predictors, arma::field& labels) { @@ -240,7 +242,7 @@ BOOST_AUTO_TEST_CASE(SortTaskTest) { task.Generate(trainPredictor, trainResponse, 8); arma::field testPredictor, testResponse; task.Generate(testPredictor, testResponse, 8); - HardCodedSortModel model; + HardCodedSortModel model(bitLen); model.Train(trainPredictor, trainResponse); arma::field predResponse; model.Predict(testPredictor, predResponse);