From 09cbc6e13aa3cb8a7c4ea6d2e1612977a40c6be7 Mon Sep 17 00:00:00 2001 From: ryan Date: Mon, 21 Dec 2015 12:32:49 -0500 Subject: [PATCH] Refactor nbc program to allow loading/saving models. --- .../naive_bayes/naive_bayes_classifier.hpp | 4 +- src/mlpack/methods/naive_bayes/nbc_main.cpp | 174 ++++++++++++------ 2 files changed, 124 insertions(+), 54 deletions(-) diff --git a/src/mlpack/methods/naive_bayes/naive_bayes_classifier.hpp b/src/mlpack/methods/naive_bayes/naive_bayes_classifier.hpp index 9647e88121..c79d836afd 100644 --- a/src/mlpack/methods/naive_bayes/naive_bayes_classifier.hpp +++ b/src/mlpack/methods/naive_bayes/naive_bayes_classifier.hpp @@ -71,8 +71,8 @@ class NaiveBayesClassifier * Train() before calling Classify(), otherwise the results may be * meaningless. */ - NaiveBayesClassifier(const size_t dimensionality, - const size_t classes); + NaiveBayesClassifier(const size_t dimensionality = 0, + const size_t classes = 0); /** * Train the Naive Bayes classifier on the given dataset. If the incremental diff --git a/src/mlpack/methods/naive_bayes/nbc_main.cpp b/src/mlpack/methods/naive_bayes/nbc_main.cpp index c573205e52..0b14fba81d 100644 --- a/src/mlpack/methods/naive_bayes/nbc_main.cpp +++ b/src/mlpack/methods/naive_bayes/nbc_main.cpp @@ -24,87 +24,157 @@ PROGRAM_INFO("Parametric Naive Bayes Classifier", "use an incremental algorithm for calculating variance. This is slower, " "but can help avoid loss of precision in some cases."); -PARAM_STRING_REQ("train_file", "A file containing the training set.", "t"); -PARAM_STRING_REQ("test_file", "A file containing the test set.", "T"); +// Model loading/saving. +PARAM_STRING("input_model_file", "File containing input Naive Bayes model.", + "m", ""); +PARAM_STRING("output_model_file", "File to save trained Naive Bayes model to.", + "M", ""); +// Training parameters. +PARAM_STRING("training_file", "A file containing the training set.", "t", ""); PARAM_STRING("labels_file", "A file containing labels for the training set.", "l", ""); -PARAM_STRING("output_file", "The file in which the predicted labels for the " - "test set will be written.", "o", "output.csv"); PARAM_FLAG("incremental_variance", "The variance of each class will be " "calculated incrementally.", "I"); +// Test parameters. +PARAM_STRING("test_file", "A file containing the test set.", "T", ""); +PARAM_STRING("output_file", "The file in which the predicted labels for the " + "test set will be written.", "o", ""); + using namespace mlpack; using namespace mlpack::naive_bayes; using namespace std; using namespace arma; +// A struct for saving the model with mappings. +struct NBCModel +{ + //! The model itself. + NaiveBayesClassifier<> nbc; + //! The mappings for labels. + Col mappings; + + //! Serialize the model. + template + void Serialize(Archive& ar, const unsigned int /* version */) + { + ar & data::CreateNVP(nbc, "nbc"); + ar & data::CreateNVP(mappings, "mappings"); + } +}; + int main(int argc, char* argv[]) { CLI::ParseCommandLine(argc, argv); // Check input parameters. - const string trainingDataFilename = CLI::GetParam("train_file"); - mat trainingData; - data::Load(trainingDataFilename, trainingData, true); + if (CLI::HasParam("training_file") && CLI::HasParam("input_model_file")) + Log::Fatal << "Cannot specify both --training_file (-t) and " + << "--input_model_file (-m)!" << endl; - // Normalize labels. - Row labels; - Col mappings; + if (!CLI::HasParam("training_file") && !CLI::HasParam("input_model_file")) + Log::Fatal << "Neither --training_file (-t) nor --input_model_file (-m) are" + << " specified!" << endl; - // Did the user pass in labels? - const string labelsFilename = CLI::GetParam("labels_file"); - if (labelsFilename != "") + if (!CLI::HasParam("training_file") && CLI::HasParam("labels_file")) + Log::Warn << "--labels_file (-l) ignored because --training_file (-t) is " + << "not specified." << endl; + if (!CLI::HasParam("training_file") && CLI::HasParam("incremental_variance")) + Log::Warn << "--incremental_variance (-I) ignored because --training_file " + << "(-t) is not specified." << endl; + + if (!CLI::HasParam("output_file") && !CLI::HasParam("output_model_file")) + Log::Warn << "Neither --output_file (-o) nor --output_model_file (-M) " + << "specified; no output will be saved!" << endl; + + if (CLI::HasParam("output_file") && !CLI::HasParam("test_file")) + Log::Warn << "--output_file (-o) ignored because no test file specified " + << "with --test_file (-T)." << endl; + + if (!CLI::HasParam("output_file") && CLI::HasParam("test_file")) + Log::Warn << "--test_file (-T) specified, but classification results will " + << "not be saved because --output_file (-o) is not specified." << endl; + + // Either we have to train a model, or load a model. + NBCModel model; + if (CLI::HasParam("training_file")) { - // Load labels. - mat rawLabels; - data::Load(labelsFilename, rawLabels, true, false); + const string trainingFile = CLI::GetParam("training_file"); + mat trainingData; + data::Load(trainingFile, trainingData, true); - // Do the labels need to be transposed? - if (rawLabels.n_cols == 1) - rawLabels = rawLabels.t(); + Row labels; - data::NormalizeLabels(rawLabels.row(0), labels, mappings); + // Did the user pass in labels? + const string labelsFilename = CLI::GetParam("labels_file"); + if (labelsFilename != "") + { + // Load labels. + mat rawLabels; + data::Load(labelsFilename, rawLabels, true, false); + + // Do the labels need to be transposed? + if (rawLabels.n_cols == 1) + rawLabels = rawLabels.t(); + + data::NormalizeLabels(rawLabels.row(0), labels, model.mappings); + } + else + { + // Use the last row of the training data as the labels. + Log::Info << "Using last dimension of training data as training labels." + << endl; + data::NormalizeLabels(trainingData.row(trainingData.n_rows - 1), labels, + model.mappings); + // Remove the label row. + trainingData.shed_row(trainingData.n_rows - 1); + } + + const bool incrementalVariance = CLI::HasParam("incremental_variance"); + + Timer::Start("nbc_training"); + model.nbc = NaiveBayesClassifier<>(trainingData, labels, + model.mappings.n_elem, incrementalVariance); + Timer::Stop("nbc_training"); } else { - // Use the last row of the training data as the labels. - Log::Info << "Using last dimension of training data as training labels." - << endl; - data::NormalizeLabels(trainingData.row(trainingData.n_rows - 1), labels, - mappings); - // Remove the label row. - trainingData.shed_row(trainingData.n_rows - 1); + // Load the model from file. + data::Load(CLI::GetParam("input_model_file"), "nbc_model", model); } - const string testingDataFilename = CLI::GetParam("test_file"); - mat testingData; - data::Load(testingDataFilename, testingData, true); + // Do we need to do testing? + if (CLI::HasParam("test_file")) + { + const string testingDataFilename = CLI::GetParam("test_file"); + mat testingData; + data::Load(testingDataFilename, testingData, true); - if (testingData.n_rows != trainingData.n_rows) - Log::Fatal << "Test data dimensionality (" << testingData.n_rows << ") " - << "must be the same as training data (" << trainingData.n_rows - << ")!" << std::endl; + if (testingData.n_rows != model.nbc.Means().n_rows) + Log::Fatal << "Test data dimensionality (" << testingData.n_rows << ") " + << "must be the same as training data (" << model.nbc.Means().n_rows + << ")!" << std::endl; - const bool incrementalVariance = CLI::HasParam("incremental_variance"); + // Time the running of the Naive Bayes Classifier. + Row results; + Timer::Start("nbc_testing"); + model.nbc.Classify(testingData, results); + Timer::Stop("nbc_testing"); - // Create and train the classifier. - Timer::Start("training"); - NaiveBayesClassifier<> nbc(trainingData, labels, mappings.n_elem, - incrementalVariance); - Timer::Stop("training"); + if (CLI::HasParam("output_file")) + { + // Un-normalize labels to prepare output. + Row rawResults; + data::RevertLabels(results, model.mappings, rawResults); - // Time the running of the Naive Bayes Classifier. - Row results; - Timer::Start("testing"); - nbc.Classify(testingData, results); - Timer::Stop("testing"); + // Output results. + const string outputFilename = CLI::GetParam("output_file"); + data::Save(outputFilename, rawResults, true); + } + } - // Un-normalize labels to prepare output. - Row rawResults; - data::RevertLabels(results, mappings, rawResults); - - // Output results. - const string outputFilename = CLI::GetParam("output_file"); - data::Save(outputFilename, rawResults, true); + if (CLI::HasParam("output_model_file")) + data::Save(CLI::GetParam("output_model_file"), "nbc_model", model, + false); }