Compile the test of Naive Bayes and Decision stump
Signed-off-by: Omar Shrit <omar@shrit.me>
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user