Compile the test of Naive Bayes and Decision stump

Signed-off-by: Omar Shrit <omar@shrit.me>
This commit is contained in:
Omar Shrit
2020-07-11 14:33:19 +02:00
parent 09e7f0de19
commit 6d2dcbe971
+81 -81
View File
@@ -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<size_t> 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<size_t> 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<arma::mat>(4, 100);
// arma::Row<size_t> 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<arma::mat>(4, 100);
arma::Row<size_t> 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<arma::mat>(3, 100);
// arma::Row<size_t> otherLabels = arma::randu<arma::Row<size_t>>(100);
// DecisionStump<> xmlDs(otherData, otherLabels, 2, 3);
arma::mat otherData = arma::randu<arma::mat>(3, 100);
arma::Row<size_t> otherLabels = arma::randu<arma::Row<size_t>>(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)