From e0436b7ffd707faed9afabeb729f8ef576e44626 Mon Sep 17 00:00:00 2001 From: houyang Date: Fri, 21 Mar 2008 23:23:32 +0000 Subject: [PATCH] --- fastlib/u/houyang/svm/smo.h | 19 ++++++++---- fastlib/u/houyang/svm/svm.h | 16 +++++----- fastlib/u/houyang/svm/svm_main.cc | 50 +++++++++++++++++++++---------- 3 files changed, 55 insertions(+), 30 deletions(-) diff --git a/fastlib/u/houyang/svm/smo.h b/fastlib/u/houyang/svm/smo.h index 5907ba223b..469771dafc 100644 --- a/fastlib/u/houyang/svm/smo.h +++ b/fastlib/u/houyang/svm/smo.h @@ -77,10 +77,6 @@ class SMO { void Train(const Dataset* dataset_in); - const Kernel& kernel() const { - return kernel_; - } - Kernel& kernel() { return kernel_; } @@ -145,6 +141,9 @@ class SMO { return kernel_cache_sign_.get(i, j) * (GetLabelSign_(i) * GetLabelSign_(j)); } + /** + * Calculate kernel values + */ void CalcKernels_() { kernel_cache_sign_.Init(n_data_, n_data_); fprintf(stderr, "Kernel Start\n"); @@ -162,7 +161,11 @@ class SMO { } }; -/* Budget SMO training for 2-classes */ +/** +* Budget SMO training for 2-classes +* +* @param: input 2-classes data matrix with labels (1,-1) in the last row +*/ template void SMO::Train(const Dataset* dataset_in) { bool examine_all = true; @@ -237,7 +240,11 @@ void SMO::Train(const Dataset* dataset_in) { } } -/* SMO training iterations */ +/** +* SMO training iterations +* +* @param: indicator: whether all the +*/ template index_t SMO::TrainIteration_(bool examine_all) { index_t num_changed = 0; diff --git a/fastlib/u/houyang/svm/svm.h b/fastlib/u/houyang/svm/svm.h index fdcb1e62dc..3d51cb1427 100644 --- a/fastlib/u/houyang/svm/svm.h +++ b/fastlib/u/houyang/svm/svm.h @@ -135,9 +135,9 @@ class SVM { void InitTrain(const Dataset& dataset, int n_classes, datanode *module); void SaveModel(String modelfilename); void LoadModel(Dataset* testset, String modelfilename); - int Classify(const Vector& vector); - void BatchClassify(Dataset* testset, String testlabelfilename); - void LoadModelBatchClassify(Dataset* testset, String modelfilename, String testlabelfilename); + int Predict(const Vector& vector); + void BatchPredict(Dataset* testset, String testlabelfilename); + void LoadModelBatchPredict(Dataset* testset, String modelfilename, String testlabelfilename); }; /** @@ -472,7 +472,7 @@ void SVM::LoadModel(Dataset* testset, String modelfilename) { * @return: a label (integer) */ template -int SVM::Classify(const Vector& datum) { +int SVM::Predict(const Vector& datum) { index_t i, j, k; ArrayList keval; keval.Init(total_num_sv_); @@ -536,7 +536,7 @@ int SVM::Classify(const Vector& datum) { * @param: file name of the testing data */ template -void SVM::BatchClassify(Dataset* testset, String testlablefilename) { +void SVM::BatchPredict(Dataset* testset, String testlablefilename) { FILE *fp = fopen(testlablefilename, "w"); if (fp == NULL) { fprintf(stderr, "Cannot save test labels to file!"); @@ -547,7 +547,7 @@ void SVM::BatchClassify(Dataset* testset, String testlablefilename) { for (index_t i = 0; i < testset->n_points(); i++) { Vector testvec; testset->matrix().MakeColumnSubvector(i, 0, num_features_, &testvec); - int testlabel = Classify(testvec); + int testlabel = Predict(testvec); if (testlabel != testset->matrix().get(num_features_, i)) err_ct++; /* save classified labels to file*/ @@ -568,9 +568,9 @@ void SVM::BatchClassify(Dataset* testset, String testlablefilename) { * @param: name of the file to store classified labels */ template -void SVM::LoadModelBatchClassify(Dataset* testset, String modelfilename, String testlabelfilename) { +void SVM::LoadModelBatchPredict(Dataset* testset, String modelfilename, String testlabelfilename) { LoadModel(testset, modelfilename); - BatchClassify(testset, testlabelfilename); + BatchPredict(testset, testlabelfilename); } #endif diff --git a/fastlib/u/houyang/svm/svm_main.cc b/fastlib/u/houyang/svm/svm_main.cc index c5d8bc03d9..23de8f437a 100644 --- a/fastlib/u/houyang/svm/svm_main.cc +++ b/fastlib/u/houyang/svm/svm_main.cc @@ -3,8 +3,10 @@ * * @file svm_main.cc * - * This file contains main routines for performing multiclass SVM - * classification. One-vs-One method is employed. + * This file contains main routines for performing multiclass + * SVM classification (SVC, one-vs-one method is employed) and + * SVM regression (epsilon-insensitive loss i.e. epsilon-SVR; and + * automatic adjustment of epsilon, i.e. mu-SVR). * * It provides four modes: * "cv": cross validation; @@ -159,7 +161,7 @@ int LoadData(Dataset* dataset, String datafilename){ } /** -* Multiclass SVM classification- Main function +* Multiclass SVM classification/ SVM regression - Main function * * @param: argc * @param: argv @@ -170,6 +172,22 @@ int main(int argc, char *argv[]) { String mode = fx_param_str_req(NULL, "mode"); String kernel = fx_param_str_req(NULL, "kernel"); + String learner_name = fx_param_str_req(NULL,"learner_name"); + int learner_typeid; + + if (learner_name == "svc") { // Support Vector Classfication + learner_typeid = 0; + } + else if (learner_name == "svr") { // Support Vector Regression + learner_typeid = 1; + } + else if (learner_name == "svde") { // Support Vector Density Estimation + learner_typeid = 2; + } + else { + fprintf(stderr, "Unknown support vector learner name! Program stops!\n"); + return 0; + } // TODO: more kernels to be supported @@ -183,20 +201,20 @@ int main(int argc, char *argv[]) { return 1; if (kernel == "linear") { - SimpleCrossValidator< SVM > cross_validator; + GeneralCrossValidator< SVM > cross_validator; /* Initialize n_folds_, confusion_matrix_; k_cv: number of cross-validation folds, need k_cv>1 */ - cross_validator.Init(&cvset,cvset.n_labels(),fx_param_int_req(NULL,"k_cv"), fx_root, "svm"); + cross_validator.Init(learner_typeid, fx_param_int_req(NULL,"k_cv"), &cvset, fx_root, "svm"); /* k_cv folds cross validation; (true): do training set permutation */ cross_validator.Run(true); - cross_validator.confusion_matrix().PrintDebug("confusion matrix"); + //cross_validator.confusion_matrix().PrintDebug("confusion matrix"); } else if (kernel == "gaussian") { - SimpleCrossValidator< SVM > cross_validator; + GeneralCrossValidator< SVM > cross_validator; /* Initialize n_folds_, confusion_matrix_; k_cv: number of cross-validation folds */ - cross_validator.Init(&cvset,cvset.n_labels(),fx_param_int_req(NULL,"k_cv"), fx_root, "svm"); + cross_validator.Init(learner_typeid, fx_param_int_req(NULL,"k_cv"), &cvset, fx_root, "svm"); /* k_cv folds cross validation; (true): do training set permutation */ cross_validator.Run(true); - cross_validator.confusion_matrix().PrintDebug("confusion matrix"); + //cross_validator.confusion_matrix().PrintDebug("confusion matrix"); } } /* Training Mode, need training data | Training + Testing(online) Mode, need training data + testing data */ @@ -216,12 +234,12 @@ int main(int argc, char *argv[]) { svm.InitTrain(trainset, trainset.n_labels(), svm_module); /* training and testing, thus no need to load model from file */ if (mode=="train_test"){ - fprintf(stderr, "SVM Classifying... \n"); + fprintf(stderr, "SVM Predicting... \n"); /* Load testing data */ Dataset testset; if (LoadData(&testset, "test_data") == 0) // TODO:param_req return 1; - svm.BatchClassify(&testset, "testlabels"); + svm.BatchPredict(&testset, "testlabels"); } } else if (kernel == "gaussian") { @@ -229,18 +247,18 @@ int main(int argc, char *argv[]) { svm.InitTrain(trainset, trainset.n_labels(), svm_module); /* training and testing, thus no need to load model from file */ if (mode=="train_test"){ - fprintf(stderr, "SVM Classifying... \n"); + fprintf(stderr, "SVM Predicting... \n"); /* Load testing data */ Dataset testset; if (LoadData(&testset, "test_data") == 0) // TODO:param_req return 1; - svm.BatchClassify(&testset, "testlabels"); // TODO:param_req + svm.BatchPredict(&testset, "testlabels"); // TODO:param_req } } } /* Testing(offline) Mode, need loading model file and testing data */ else if (mode=="test") { - fprintf(stderr, "SVM Classifying... \n"); + fprintf(stderr, "SVM Predicting... \n"); /* Load testing data */ Dataset testset; @@ -253,12 +271,12 @@ int main(int argc, char *argv[]) { if (kernel == "linear") { SVM svm; svm.Init(testset, testset.n_labels(), svm_module); // TODO:n_labels() -> num_classes_ - svm.LoadModelBatchClassify(&testset, "svm_model", "testlabels"); // TODO:param_req + svm.LoadModelBatchPredict(&testset, "svm_model", "testlabels"); // TODO:param_req } else if (kernel == "gaussian") { SVM svm; svm.Init(testset, testset.n_labels(), svm_module); // TODO:n_labels() -> num_classes_ - svm.LoadModelBatchClassify(&testset, "svm_model", "testlabels"); // TODO:param_req + svm.LoadModelBatchPredict(&testset, "svm_model", "testlabels"); // TODO:param_req } } fx_done();