In GetLabels(), all the input parameters are required to init beforehand. GeneralCrossValidator() is modified accordingly.

This commit is contained in:
houyang
2008-04-22 14:37:56 +00:00
parent c6f028e4f5
commit 610e76341b
3 changed files with 16 additions and 4 deletions
+9 -4
View File
@@ -559,6 +559,11 @@ void GeneralCrossValidator<TLearner>::Run(bool randomized) {
ArrayList<index_t> cv_labels_startpos;
// Get label list and label indices from the cross validation data set
index_t num_classes = data_->n_labels();
cv_labels_list.Init();
cv_labels_index.Init();
cv_labels_ct.Init();
cv_labels_startpos.Init();
data_->GetLabels(cv_labels_list, cv_labels_index, cv_labels_ct, cv_labels_startpos);
// randomize the original data set within each class if necessary
@@ -602,7 +607,7 @@ void GeneralCrossValidator<TLearner>::Run(bool randomized) {
VERBOSE_MSG(1, "cross: Training fold %d", i_fold);
fx_timer_start(foldmodule, "train");
// training
classifier.InitTrain(learner_typeid_, train, clsf_n_classes_, learner_module);
classifier.InitTrain(learner_typeid_, train, learner_module);
fx_timer_stop(foldmodule, "train");
// validation; measure method: percent of correctly classified validation samples
@@ -616,7 +621,7 @@ void GeneralCrossValidator<TLearner>::Run(bool randomized) {
validation.matrix().MakeColumnVector(i, &validation_vector_with_label);
validation_vector_with_label.MakeSubvector(0, validation.n_features()-1, &validation_vector);
// testing (classification)
int label_predict = int(classifier.Predict(validation_vector));
int label_predict = int(classifier.Predict(learner_typeid_, validation_vector));
double label_expect_dbl = validation_vector_with_label[validation.n_features()-1];
int label_expect = int(label_expect_dbl);
@@ -683,7 +688,7 @@ void GeneralCrossValidator<TLearner>::Run(bool randomized) {
VERBOSE_MSG(1, "cross: Training fold %d", i_fold);
fx_timer_start(foldmodule, "train");
// training
learner.InitTrain(learner_typeid_, train, 0, learner_module); // 0: dummy number of classes
learner.InitTrain(learner_typeid_, train, learner_module); // 0: dummy number of classes
fx_timer_stop(foldmodule, "train");
// validation
@@ -697,7 +702,7 @@ void GeneralCrossValidator<TLearner>::Run(bool randomized) {
validation_vector_with_label.MakeSubvector(
0, validation.n_features()-1, &validation_vector);
// testing
double value_predict = learner.Predict(validation_vector);
double value_predict = learner.Predict(learner_typeid_, validation_vector);
double value_true = validation_vector_with_label[validation.n_features()-1];
double value_err = value_predict - value_true;
+5
View File
@@ -295,6 +295,11 @@ void Dataset::GetLabels(ArrayList<double> &labels_list,
index_t n_points = matrix_.n_cols();
index_t n_labels = 0;
labels_list.Destruct();
labels_index.Destruct();
labels_ct.Destruct();
labels_startpos.Destruct();
labels_index.Init(n_points);
labels_list.Init();
labels_ct.Init();
+2
View File
@@ -402,6 +402,8 @@ class Dataset {
* class_2...class_k), each item indicate the position of the label
* in the dataset.
*
* All input parameters need to be initilized beforehand.
*
* @param labels_list a list of labels in the dataset. e.g. [0.0,1.0,2.0]
* for a 3-class dataset
* @param labels_index the label indices of each data point. e.g.