This commit is contained in:
houyang
2008-03-21 23:23:32 +00:00
parent f849e3f7cd
commit e0436b7ffd
3 changed files with 55 additions and 30 deletions
+13 -6
View File
@@ -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<typename TKernel>
void SMO<TKernel>::Train(const Dataset* dataset_in) {
bool examine_all = true;
@@ -237,7 +240,11 @@ void SMO<TKernel>::Train(const Dataset* dataset_in) {
}
}
/* SMO training iterations */
/**
* SMO training iterations
*
* @param: indicator: whether all the
*/
template<typename TKernel>
index_t SMO<TKernel>::TrainIteration_(bool examine_all) {
index_t num_changed = 0;
+8 -8
View File
@@ -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<TKernel>::LoadModel(Dataset* testset, String modelfilename) {
* @return: a label (integer)
*/
template<typename TKernel>
int SVM<TKernel>::Classify(const Vector& datum) {
int SVM<TKernel>::Predict(const Vector& datum) {
index_t i, j, k;
ArrayList<double> keval;
keval.Init(total_num_sv_);
@@ -536,7 +536,7 @@ int SVM<TKernel>::Classify(const Vector& datum) {
* @param: file name of the testing data
*/
template<typename TKernel>
void SVM<TKernel>::BatchClassify(Dataset* testset, String testlablefilename) {
void SVM<TKernel>::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<TKernel>::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<TKernel>::BatchClassify(Dataset* testset, String testlablefilename) {
* @param: name of the file to store classified labels
*/
template<typename TKernel>
void SVM<TKernel>::LoadModelBatchClassify(Dataset* testset, String modelfilename, String testlabelfilename) {
void SVM<TKernel>::LoadModelBatchPredict(Dataset* testset, String modelfilename, String testlabelfilename) {
LoadModel(testset, modelfilename);
BatchClassify(testset, testlabelfilename);
BatchPredict(testset, testlabelfilename);
}
#endif
+34 -16
View File
@@ -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<SVMLinearKernel> > cross_validator;
GeneralCrossValidator< SVM<SVMLinearKernel> > 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<SVMRBFKernel> > cross_validator;
GeneralCrossValidator< SVM<SVMRBFKernel> > 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<SVMLinearKernel> 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<SVMRBFKernel> 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();