diff --git a/src/mlpack/tests/serialization_test.cpp b/src/mlpack/tests/serialization_test.cpp index 8ca7b00ebd..d8f54ca9ae 100644 --- a/src/mlpack/tests/serialization_test.cpp +++ b/src/mlpack/tests/serialization_test.cpp @@ -891,63 +891,63 @@ BOOST_AUTO_TEST_CASE(SoftmaxRegressionTest) // } // } -// BOOST_AUTO_TEST_CASE(NaiveBayesSerializationTest) -// { -// // Train NBC randomly. Make sure the model is the same after serializing and -// // re-loading. -// arma::mat dataset; -// dataset.randu(10, 500); -// arma::Row labels(500); -// for (size_t i = 0; i < 500; ++i) -// { -// if (dataset(0, i) > 0.5) -// labels[i] = 0; -// else -// labels[i] = 1; -// } +BOOST_AUTO_TEST_CASE(NaiveBayesSerializationTest) +{ + // Train NBC randomly. Make sure the model is the same after serializing and + // re-loading. + arma::mat dataset; + dataset.randu(10, 500); + arma::Row labels(500); + for (size_t i = 0; i < 500; ++i) + { + if (dataset(0, i) > 0.5) + labels[i] = 0; + else + labels[i] = 1; + } -// NaiveBayesClassifier<> nbc(dataset, labels, 2); + NaiveBayesClassifier<> nbc(dataset, labels, 2); -// // Initialize some empty Naive Bayes classifiers. -// NaiveBayesClassifier<> xmlNbc(0, 0), textNbc(0, 0), binaryNbc(0, 0); -// SerializeObjectAll(nbc, xmlNbc, textNbc, binaryNbc); + // Initialize some empty Naive Bayes classifiers. + NaiveBayesClassifier<> xmlNbc(0, 0), textNbc(0, 0), binaryNbc(0, 0); + SerializeObjectAll(nbc, xmlNbc, textNbc, binaryNbc); -// BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, xmlNbc.Means().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, textNbc.Means().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, binaryNbc.Means().n_elem); -// for (size_t i = 0; i < nbc.Means().n_elem; ++i) -// { -// BOOST_REQUIRE_CLOSE(nbc.Means()[i], xmlNbc.Means()[i], 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Means()[i], textNbc.Means()[i], 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Means()[i], binaryNbc.Means()[i], 1e-5); -// } + BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, xmlNbc.Means().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, textNbc.Means().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Means().n_elem, binaryNbc.Means().n_elem); + for (size_t i = 0; i < nbc.Means().n_elem; ++i) + { + BOOST_REQUIRE_CLOSE(nbc.Means()[i], xmlNbc.Means()[i], 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Means()[i], textNbc.Means()[i], 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Means()[i], binaryNbc.Means()[i], 1e-5); + } -// BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, xmlNbc.Variances().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, textNbc.Variances().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, binaryNbc.Variances().n_elem); -// for (size_t i = 0; i < nbc.Variances().n_elem; ++i) -// { -// BOOST_REQUIRE_CLOSE(nbc.Variances()[i], xmlNbc.Variances()[i], 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Variances()[i], textNbc.Variances()[i], 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Variances()[i], binaryNbc.Variances()[i], 1e-5); -// } + BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, xmlNbc.Variances().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, textNbc.Variances().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Variances().n_elem, binaryNbc.Variances().n_elem); + for (size_t i = 0; i < nbc.Variances().n_elem; ++i) + { + BOOST_REQUIRE_CLOSE(nbc.Variances()[i], xmlNbc.Variances()[i], 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Variances()[i], textNbc.Variances()[i], 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Variances()[i], binaryNbc.Variances()[i], 1e-5); + } -// BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, -// xmlNbc.Probabilities().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, -// textNbc.Probabilities().n_elem); -// BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, -// binaryNbc.Probabilities().n_elem); -// for (size_t i = 0; i < nbc.Probabilities().n_elem; ++i) -// { -// BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], xmlNbc.Probabilities()[i], -// 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], textNbc.Probabilities()[i], -// 1e-5); -// BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], binaryNbc.Probabilities()[i], -// 1e-5); -// } -// } + BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, + xmlNbc.Probabilities().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, + textNbc.Probabilities().n_elem); + BOOST_REQUIRE_EQUAL(nbc.Probabilities().n_elem, + binaryNbc.Probabilities().n_elem); + for (size_t i = 0; i < nbc.Probabilities().n_elem; ++i) + { + BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], xmlNbc.Probabilities()[i], + 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], textNbc.Probabilities()[i], + 1e-5); + BOOST_REQUIRE_CLOSE(nbc.Probabilities()[i], binaryNbc.Probabilities()[i], + 1e-5); + } +} // BOOST_AUTO_TEST_CASE(RASearchTest) // { @@ -1065,41 +1065,41 @@ BOOST_AUTO_TEST_CASE(SoftmaxRegressionTest) // textLsh.SecondHashTable()[i], binaryLsh.SecondHashTable()[i]); // } -// // Make sure serialization works for the decision stump. -// BOOST_AUTO_TEST_CASE(DecisionStumpTest) -// { -// // Generate dataset. -// arma::mat trainingData = arma::randu(4, 100); -// arma::Row labels(100); -// for (size_t i = 0; i < 25; ++i) -// labels[i] = 0; -// for (size_t i = 25; i < 50; ++i) -// labels[i] = 3; -// for (size_t i = 50; i < 75; ++i) -// labels[i] = 1; -// for (size_t i = 75; i < 100; ++i) -// labels[i] = 2; +// Make sure serialization works for the decision stump. +BOOST_AUTO_TEST_CASE(DecisionStumpTest) +{ + // Generate dataset. + arma::mat trainingData = arma::randu(4, 100); + arma::Row labels(100); + for (size_t i = 0; i < 25; ++i) + labels[i] = 0; + for (size_t i = 25; i < 50; ++i) + labels[i] = 3; + for (size_t i = 50; i < 75; ++i) + labels[i] = 1; + for (size_t i = 75; i < 100; ++i) + labels[i] = 2; -// DecisionStump<> ds(trainingData, labels, 4, 3); + DecisionStump<> ds(trainingData, labels, 4, 3); -// arma::mat otherData = arma::randu(3, 100); -// arma::Row otherLabels = arma::randu>(100); -// DecisionStump<> xmlDs(otherData, otherLabels, 2, 3); + arma::mat otherData = arma::randu(3, 100); + arma::Row otherLabels = arma::randu>(100); + DecisionStump<> xmlDs(otherData, otherLabels, 2, 3); -// DecisionStump<> textDs; -// DecisionStump<> binaryDs(trainingData, labels, 4, 10); + DecisionStump<> textDs; + DecisionStump<> binaryDs(trainingData, labels, 4, 10); -// SerializeObjectAll(ds, xmlDs, textDs, binaryDs); + SerializeObjectAll(ds, xmlDs, textDs, binaryDs); -// // Make sure that everything is the same about the new decision stumps. -// BOOST_REQUIRE_EQUAL(ds.SplitDimension(), xmlDs.SplitDimension()); -// BOOST_REQUIRE_EQUAL(ds.SplitDimension(), textDs.SplitDimension()); -// BOOST_REQUIRE_EQUAL(ds.SplitDimension(), binaryDs.SplitDimension()); + // Make sure that everything is the same about the new decision stumps. + BOOST_REQUIRE_EQUAL(ds.SplitDimension(), xmlDs.SplitDimension()); + BOOST_REQUIRE_EQUAL(ds.SplitDimension(), textDs.SplitDimension()); + BOOST_REQUIRE_EQUAL(ds.SplitDimension(), binaryDs.SplitDimension()); -// CheckMatrices(ds.Split(), xmlDs.Split(), textDs.Split(), binaryDs.Split()); -// CheckMatrices(ds.BinLabels(), xmlDs.BinLabels(), textDs.BinLabels(), -// binaryDs.BinLabels()); -// } + CheckMatrices(ds.Split(), xmlDs.Split(), textDs.Split(), binaryDs.Split()); + CheckMatrices(ds.BinLabels(), xmlDs.BinLabels(), textDs.BinLabels(), + binaryDs.BinLabels()); +} // // Make sure serialization works for LARS. // BOOST_AUTO_TEST_CASE(LARSTest)