From 83c95e2aa868c587db2d3c595026a68fa0e8e7ef Mon Sep 17 00:00:00 2001 From: Garry Boyer Date: Thu, 5 Apr 2007 07:52:44 +0000 Subject: [PATCH] linear svm seems to be working --- fastlib/u/garryb/svm/smo.h | 100 +++++++++++++++++++++--------------- fastlib/u/garryb/svm/svm.cc | 29 +++++++++-- fastlib/u/garryb/svm/svm.h | 13 ++++- 3 files changed, 96 insertions(+), 46 deletions(-) diff --git a/fastlib/u/garryb/svm/smo.h b/fastlib/u/garryb/svm/smo.h index 65f859d780..dc8873debc 100644 --- a/fastlib/u/garryb/svm/smo.h +++ b/fastlib/u/garryb/svm/smo.h @@ -81,11 +81,13 @@ class SMO { } bool IsBound_(double alpha) const { - return alpha <= 0 || alpha >= c_; + return alpha == 0 || alpha == c_; } double GetLabelSign_(index_t i) const { - return matrix_.get(matrix_.n_rows()-1, i) * 2.0 - 1.0; + double v = matrix_.get(matrix_.n_rows()-1, i) * 2.0 - 1.0; + //DEBUG_MSG(0, "v = %f", v); + return v; } void GetVector_(index_t i, Vector *v) const { @@ -93,11 +95,16 @@ class SMO { } double Error_(index_t i) const { + double val; if (!IsBound_(alpha_[i])) { - return error_[i]; + val = error_[i]; +#ifdef VERBOSE + DEBUG_MSG(0, "error values %f and %f", error_[i], Evaluate_(i) - GetLabelSign_(i)); +#endif } else { - return Evaluate_(i) - GetLabelSign_(i); + val = Evaluate_(i) - GetLabelSign_(i); } + return val; } double Evaluate_(index_t i) const; @@ -137,6 +144,8 @@ void SMO::GetSVM(Matrix *support_vectors, Vector *support_alpha) const dest.CopyValues(source); (*support_alpha)[i_support] = alpha_[i] * GetLabelSign_(i); + + i_support++; } } } @@ -149,7 +158,7 @@ double SMO::Evaluate_(index_t i) const { double summation = 0; // TODO: This is linear in the size of the training points - for (index_t j = 0; j < matrix_.n_cols(); j++) { + for (index_t j = 0; j < alpha_.length(); j++) { if (likely(alpha_[j] != 0)) { Vector support_vector; GetVector_(j, &support_vector); @@ -161,23 +170,24 @@ double SMO::Evaluate_(index_t i) const { } } - return summation - thresh_; + return (summation - thresh_); } template void SMO::Train() { bool examine_all = true; + index_t num_changed = 0; - do { + while (num_changed > 0 || examine_all) { DEBUG_GOT_HERE(0); - index_t num_changed = TrainIteration_(examine_all); + num_changed = TrainIteration_(examine_all); if (examine_all) { examine_all = false; } else if (num_changed == 0) { examine_all = true; } - } while (examine_all); + } } template @@ -185,7 +195,7 @@ index_t SMO::TrainIteration_(bool examine_all) { index_t num_changed = 0; for (index_t i = 0; i < alpha_.length(); i++) { - if ((examine_all || IsBound_(alpha_[i])) && TryChange_(i)) { + if ((examine_all || !IsBound_(alpha_[i])) && TryChange_(i)) { num_changed++; } } @@ -198,6 +208,8 @@ bool SMO::TryChange_(index_t j) { double error_j = Error_(j); // WALDO double rj = error_j * GetLabelSign_(j); + DEBUG_GOT_HERE(0); + if (!((rj < -SMO_TOLERANCE && alpha_[j] < c_) || (rj > SMO_TOLERANCE && alpha_[j] > 0))) { return false; // nothing changed @@ -205,27 +217,26 @@ bool SMO::TryChange_(index_t j) { // first try the one we suspect to have the largest yield - if (error_j > 0) { + if (error_j != 0) { index_t i = -1; double error_i = error_j; - for (index_t k = 0; k < alpha_.length(); k++) { - if (!IsBound_(alpha_[k]) && error_[k] < error_i) { - error_i = error_[k]; - i = k; - } - } - if (i != -1 && TakeStep_(i, j, error_j)) { - return true; - } - } else if (likely(error_j < 0)) { - index_t i = -1; - double error_i = error_j; - for (index_t k = 0; k < alpha_.length(); k++) { - if (!IsBound_(alpha_[k]) && error_[k] > error_i) { - error_i = error_[k]; - i = k; + + if (error_j > 0) { + for (index_t k = 0; k < alpha_.length(); k++) { + if (!IsBound_(alpha_[k]) && error_[k] < error_i) { + error_i = error_[k]; + i = k; + } + } + } else { + for (index_t k = 0; k < alpha_.length(); k++) { + if (!IsBound_(alpha_[k]) && error_[k] > error_i) { + error_i = error_[k]; + i = k; + } } } + if (i != -1 && TakeStep_(i, j, error_j)) { return true; } @@ -261,14 +272,15 @@ bool SMO::TryChange_(index_t j) { template bool SMO::TakeStep_(index_t i, index_t j, double error_j) { if (i == j) { + DEBUG_GOT_HERE(0); return false; } double yi = GetLabelSign_(i); - double yj = GetLabelSign_(i); + double yj = GetLabelSign_(j); double alpha_i; double alpha_j; - double thresh_new; + double d_thresh; double l; double u; double s = yi * yj; @@ -286,6 +298,8 @@ bool SMO::TakeStep_(index_t i, index_t j, double error_j) { if (l == u) { // TODO: might put in some tolerance + DEBUG_MSG(0, "l=%f, u=%f, r=%f, c_=%f, s=%f", l, u, r, c_, s); + DEBUG_GOT_HERE(0); return false; } @@ -294,13 +308,16 @@ bool SMO::TakeStep_(index_t i, index_t j, double error_j) { double kij = EvalKernel_(i, j); double kjj = EvalKernel_(j, j); // second derivative of objective function - double eta = 2 * kij - kii - kjj; + double eta = +2*kij - kii - kjj; + DEBUG_MSG(0, "kij=%f, kii=%f, kjj=%f", kij, kii, kjj); if (likely(eta < 0)) { + DEBUG_MSG(0, "Common case"); alpha_j = alpha_[j] - yj * (error_i - error_j) / eta; alpha_j = math::ClampRange(alpha_j, l, u); } else { + DEBUG_MSG(0, "Uncommon case"); double fiold = error_i + yi; double fjold = error_j + yj; double vi = fiold + thresh_ - yi*alpha_[i]*kii - yj*alpha_[j]*kij; @@ -329,6 +346,7 @@ bool SMO::TakeStep_(index_t i, index_t j, double error_j) { // check if there is progress if (fabs(d_alpha_j) < SMO_EPS*(alpha_j + alpha_[j] + SMO_EPS)) { + DEBUG_GOT_HERE(0); return false; } @@ -336,15 +354,15 @@ bool SMO::TakeStep_(index_t i, index_t j, double error_j) { double d_alpha_i = alpha_i - alpha_[i]; // calculate threshold - double thresh_i = thresh_ + error_i + yi*d_alpha_i*kii + yj*d_alpha_j*kij; - double thresh_j = thresh_ + error_j + yi*d_alpha_i*kij + yj*d_alpha_j*kjj; + double d_thresh_i = error_i + yi*d_alpha_i*kii + yj*d_alpha_j*kij; + double d_thresh_j = error_j + yi*d_alpha_i*kij + yj*d_alpha_j*kjj; if (!IsBound_(alpha_i)) { - thresh_new = thresh_i; + d_thresh = d_thresh_i; } else if (!IsBound_(alpha_j)) { - thresh_new = thresh_j; + d_thresh = d_thresh_j; } else { - thresh_new = (thresh_i + thresh_j) / 2.0; + d_thresh = (d_thresh_i + d_thresh_j) / 2.0; } // if not bound, error must be zero @@ -355,23 +373,23 @@ bool SMO::TakeStep_(index_t i, index_t j, double error_j) { error_[j] = 0; } if (!IsBound_(alpha_i) && !IsBound_(alpha_j)) { - fprintf(stderr, "Neither ai nor aj are bound."); + DEBUG_MSG(0, "Neither ai nor aj are bound."); } double ti = yi*d_alpha_i; - double tj = yi*d_alpha_j; - double d_thresh = thresh_new - thresh_; + double tj = yj*d_alpha_j; for (index_t k = 0; k < error_.length(); k++) { - if (likely(k != i)) { + if (likely(k != i) && likely(k != j) && !IsBound_(alpha_[k])) { error_[k] += ti*EvalKernel_(i, k) + tj*EvalKernel_(j, k) - d_thresh; } } - thresh_ = thresh_new; + thresh_ += d_thresh; alpha_[i] = alpha_i; alpha_[j] = alpha_j; - + + DEBUG_GOT_HERE(0); return true; } diff --git a/fastlib/u/garryb/svm/svm.cc b/fastlib/u/garryb/svm/svm.cc index 02a23edb94..7f67233922 100644 --- a/fastlib/u/garryb/svm/svm.cc +++ b/fastlib/u/garryb/svm/svm.cc @@ -5,10 +5,33 @@ int main(int argc, char *argv[]) { Dataset dataset; - if (!PASSED(dataset.InitFromFile(fx_param_str_req(NULL, "data")))) { - fprintf(stderr, "Couldn't open the data file."); - return 1; + //if (!PASSED(dataset.InitFromFile(fx_param_str_req(NULL, "data")))) { + // fprintf(stderr, "Couldn't open the data file."); + // return 1; + //} + Matrix m; + index_t n = fx_param_int(NULL, "n", 30); + double slope = fx_param_double(NULL, "slope", 1.0); + double margin = fx_param_double(NULL, "margin", 1.0); + double var = fx_param_double(NULL, "var", 1.0); + + m.Init(3, n); + + for (index_t i = 0; i < n; i += 2) { + double x = (rand() * 2.0 / RAND_MAX) - 1.0; + double y = margin / 2 + (rand() * var / RAND_MAX); + m.set(0, i, x); + m.set(1, i, x*slope+y); + m.set(2, i, 0); + m.set(0, i+1, x); + m.set(1, i+1, x*slope-y); + m.set(2, i+1, 1); } + //Matrix m2; + //la::TransposeInit(m, &m2); + //m2.PrintDebug("m"); + + dataset.AliasMatrix(m); SimpleCrossValidator< SVM > cross_validator; cross_validator.Init(&dataset, 2, 4, fx_root, "svm"); diff --git a/fastlib/u/garryb/svm/svm.h b/fastlib/u/garryb/svm/svm.h index 1eddbfdb04..5f029b098b 100644 --- a/fastlib/u/garryb/svm/svm.h +++ b/fastlib/u/garryb/svm/svm.h @@ -44,7 +44,7 @@ void SVM::InitTrain( kernel_.Init(fx_submodule(module, "kernel", "kernel")); - c_ = fx_param_double(module, "c", 0.01); + c_ = fx_param_double(module, "c", 1.0); SMO smo; smo.Init(&dataset, c_); @@ -55,6 +55,11 @@ void SVM::InitTrain( smo.GetSVM(&support_vectors_, &alpha_); DEBUG_ASSERT(alpha_.length() != 0); DEBUG_ASSERT(alpha_.length() == support_vectors_.n_cols()); + + DEBUG_ONLY(fprintf(stderr, "----------------------\n")); + DEBUG_ONLY(support_vectors_.PrintDebug("support vectors")); + DEBUG_ONLY(alpha_.PrintDebug("support vector weights")); + DEBUG_ONLY(fprintf(stderr, "-- THRESHOLD: %f\n", thresh_)); } template @@ -64,10 +69,14 @@ int SVM::Classify(const Vector& datum) { for (index_t i = 0; i < alpha_.length(); i++) { Vector support_vector; support_vectors_.MakeColumnVector(i, &support_vector); + double term = alpha_[i] * kernel_.Eval(datum, support_vector); - summation += alpha_[i] * kernel_.Eval(datum, support_vector); + DEBUG_MSG(0, "alpha %f, term %f", alpha_[i], term); + + summation += term; } + DEBUG_MSG(0, "summation=%f, thresh_=%f", summation, thresh_); return (summation - thresh_ > 0.0) ? 1 : 0; }