アットウィキロゴ

sd.cpp sd.hpp

#ifndef __SD_H__
#define __SD_H__

#include "env.hpp"

#define ARMIJO_MU 0.01
#define ARMIJO_MAX 30

#define MOMENTUM_ALPHA 0.9
#define REGULAR_LAMBDA 1.0

class Sd : public Env{
public:
  Sd(
      std::string _data,
      unsigned _hid,
      std::string _projName = "testSd",
      double _min = -1.0,
      double _max = 1.0
  );
  ~Sd();

  void show();

  bool learn(
     unsigned _times,
     double _eta,
     bool _reset = true
  );

  bool saveEta();

protected:
  void setup();
  bool setGSLVector();
  void releaseGSLVector();

  double bprop( gsl_vector *x, gsl_vector *t );
  double bbprop();
  void renewWeight();
  void armijo( double _mu = ARMIJO_MU );
  void bpq();
  double calcError();
  double calcEval();

  unsigned times;
  double eta;
  double alpha;
  double lambda;

  std::vector<double> logEta;

  gsl_vector *g;
  gsl_vector *d;
  gsl_vector *dw;
  gsl_vector *dg;

  // temp for bprop
  gsl_vector *d1;
  gsl_vector *d2;
  gsl_vector_view view_gth1;
  gsl_vector_view view_gth2;
  gsl_matrix_view view_gw1;
  gsl_matrix_view view_gw2;

  // temp for bbprop
  gsl_vector *sum_g;

  //temp for armijo
  gsl_vector *try_d;

  // temp for bpq
  gsl_vector *dTH;
  gsl_vector *ha1;
  gsl_vector *hz;
  gsl_vector *ha2;
  gsl_vector *hy;
  gsl_vector *hd1;
  gsl_vector *hd2;
  gsl_vector *hsig;
  gsl_vector *sum_dTH;
  gsl_vector_view view_dth1;
  gsl_vector_view view_dth2;
  gsl_matrix_view view_dw1;
  gsl_matrix_view view_dw2;
  gsl_vector_view view_dTHth1;
  gsl_vector_view view_dTHth2;
  gsl_matrix_view view_dTHw1;
  gsl_matrix_view view_dTHw2;

};

#endif /* __SD_H__ */


///////////////////////////////////////////////////////////////////////////////

using namespace std;

Sd::Sd(string _data, unsigned _hid, string _projName, double _min, double _max)
  : Env(_data, _hid, _projName, _min, _max)
{
  setup();
  setGSLVector();
  cout << "Learn by SD" << endl;
}

Sd::~Sd()
{
  releaseGSLVector();
}

void Sd::setup()
{
  times = 0;
  eta = 0.01;

#ifdef _MOMENTUM_
  alpha = MOMENTUM_ALPHA;
#else
  alpha = 0.0;
#endif /* set alpha */

#ifdef _REGULAR_
  lambda = REGULAR_LAMBDA;
#else
  lambda = 0.0;
#endif /* set lambda */

}

bool Sd::setGSLVector()
{
  g = gsl_vector_calloc( num_total );
  d = gsl_vector_calloc( num_total );
  dw = gsl_vector_calloc( num_total );
  dg = gsl_vector_calloc( num_total );

  d1 = gsl_vector_calloc( num_hid );
  d2 = gsl_vector_calloc( num_out );
  view_gth1 = gsl_vector_subvector(g, idx_th1, len_th1);
  view_gth2 = gsl_vector_subvector(g, idx_th2, len_th2);
  view_gw1 = gsl_matrix_view_array(&g->data[idx_w1], num_hid, num_in);
  view_gw2 = gsl_matrix_view_array(&g->data[idx_w2], num_out, num_hid);

  sum_g = gsl_vector_calloc( num_total );

  try_d = gsl_vector_calloc( num_total );

  dTH = gsl_vector_calloc( num_total );
  ha1 = gsl_vector_calloc( num_hid );
  hz = gsl_vector_calloc( num_hid );
  ha2 = gsl_vector_calloc( num_out );
  hy = gsl_vector_calloc( num_out );
  hd1 = gsl_vector_calloc( num_hid );
  hd2 = gsl_vector_calloc( num_out );
  hsig = gsl_vector_calloc( num_hid );
  sum_dTH = gsl_vector_calloc( num_total );
  view_dth1 = gsl_vector_subvector(d, idx_th1, len_th1);
  view_dth2 = gsl_vector_subvector(d, idx_th2, len_th2);
  view_dw1 = gsl_matrix_view_array(&d->data[idx_w1], num_hid, num_in);
  view_dw2 = gsl_matrix_view_array(&d->data[idx_w2], num_out, num_hid);
  view_dTHth1 = gsl_vector_subvector(dTH, idx_th1, len_th1);
  view_dTHth2 = gsl_vector_subvector(dTH, idx_th2, len_th2);
  view_dTHw1 = gsl_matrix_view_array(&dTH->data[idx_w1], num_hid, num_in);
  view_dTHw2 = gsl_matrix_view_array(&dTH->data[idx_w2], num_out, num_hid);

  return true;
}

void Sd::releaseGSLVector()
{
  gsl_vector_free(g);
  gsl_vector_free(d);
  gsl_vector_free(dw);
  gsl_vector_free(dg);

  gsl_vector_free(d1);
  gsl_vector_free(d2);

  gsl_vector_free(sum_g);

  gsl_vector_free(try_d);

  gsl_vector_free(dTH);
  gsl_vector_free(ha1);
  gsl_vector_free(hz);
  gsl_vector_free(ha2);
  gsl_vector_free(hy);
  gsl_vector_free(hd1);
  gsl_vector_free(hd2);
  gsl_vector_free(hsig);
  gsl_vector_free(sum_dTH);

  cout << "Sd::releaseGSLVector() succeeded." << endl;
}

void Sd::show()
{
  NN::show();
  cout << endl;
  cout << "momentum alpha = " << alpha << endl;
  cout << "regularization lambda = " << lambda << endl;

#ifdef DEBUG_SD
  cout << "d1" << endl;
  show_vec(d1);
  cout << "d2" << endl;
  show_vec(d2);
#endif
}

// calc g of 1-data without regularization
double Sd::bprop(gsl_vector *x, gsl_vector *t)
{
  // FP : set a1 a2 z y
  fprop(x);

  // calculate output layers
  gsl_vector_memcpy(d2, y); // d2 := y
  gsl_vector_sub(d2, t); // d2 := d2 - out = y - out
  
  // calculate hidden layers
  gsl_blas_dgemv(CblasTrans, 1.0, w2, d2, 0.0, d1);  // d1 := d2 w2 = w2T d2
  apply_vec(a1, f_act_d, a1); // a1 changed here!!!
  gsl_vector_mul(d1, a1);
  
  // calculate g
  gsl_vector_memcpy( &view_gth1.vector, d1 );
  gsl_vector_memcpy( &view_gth2.vector, d2 );

  gsl_matrix_set_zero( &view_gw1.matrix );
  gsl_blas_dger( 1.0, d1, x, &view_gw1.matrix );

  gsl_matrix_set_zero( &view_gw2.matrix );
  gsl_blas_dger( 1.0, d2, z, &view_gw2.matrix );

  //////////////////////
  //check NA of w1 and w2 here.
  //////////////////////
  
  // return error
  return sq_vec( d2 );
}

// calc g of all data with regularization
double Sd::bbprop()
{
  gsl_vector_set_zero( sum_g );
  double err = 0.0;

  for(unsigned n = 0; n < num_data; n++){
    err += bprop(x[n], t[n]);
    gsl_vector_add( sum_g, g);
  }
  gsl_vector_memcpy(g, sum_g);

#ifdef _REGULAR_
  add_vec(g, lambda, w);
#endif

  return err;
}

void Sd::renewWeight()
{
  gsl_vector_scale(dw, alpha);
  add_vec(dw, eta, d);
  gsl_vector_add(w, dw);
}

// it will break a1 a2 z y, d2
double Sd::calcError()
{
  double err = 0.0;
  // calc error
  for(unsigned n = 0; n < num_data; n++){
    // FP : set a1 a2 z y
    fprop( x[n] );

    // calc output layers
    gsl_vector_memcpy(d2, y); // d2 := y
    gsl_vector_sub(d2, t[n]); // d2 := d2 - out = y - out

    // sum error
    err += sq_vec( d2 );
  }

  return err;
}

double Sd::calcEval()
{
  double eval = calcError();
#ifdef _REGULAR_
  eval += lambda * sq_vec( w );
#endif

  return eval;
}

void Sd::armijo( double _mu )
{
  double err = calcEval();
  double err_d;

  //calc error_derivatibe
  gsl_blas_ddot(g, d, &err_d);
  err_d *= _mu;

  // suppose w(t+1) = w(t) + 1.0 * d
  gsl_vector_memcpy( try_d, d );
  gsl_vector_add( w, try_d );

  // backtrack eta and renew weights
  unsigned count = 0;
  eta = 1.0;
  while( calcEval() > err + err_d && ++count < ARMIJO_MAX ){
    gsl_vector_scale( try_d, 0.5 );
    gsl_vector_sub( w, try_d );
    err_d *= 0.5;
    eta *= 0.5;
  }
}

void Sd::bpq()
{
  double denom, numer;

  // denominator
  gsl_blas_ddot(d, g, &denom);
  denom *= -1.0;

  // numerator

  // calc dT H
  gsl_vector_set_zero( sum_dTH );
  for(unsigned n = 0; n < num_data; n++){
    fprop(x[n]);

    gsl_vector_memcpy(ha1, &view_dth1.vector);
    gsl_blas_dgemv(CblasNoTrans, 1.0, &view_dw1.matrix, x[n], 1.0, ha1);
    apply_vec(hz, f_act_d, ha1);
    gsl_vector_mul(hz, ha1);

    // calculate output layers
    gsl_vector_memcpy(ha2, &view_dth2.vector);
    gsl_blas_dgemv(CblasNoTrans, 1.0, &view_dw2.matrix, z, 1.0, ha2);
    gsl_blas_dgemv(CblasNoTrans, 1.0, w2, hz, 1.0, ha2);
    gsl_vector_memcpy(hy, ha2); // apply_vec(hy, f_out, ha2) ?

    // calculate output layers
    gsl_vector_memcpy(hd2, hy);

    // calculate hidden layers
    gsl_blas_dgemv(CblasTrans, 1.0, w2, d2, 0.0, hd1);  // d1 := d2 w2 = w2T d2
    gsl_vector_mul(hd1, ha1);
    apply_vec(hsig, f_act_d, ha1);

      // calc 1 - sig(ha1)
      apply_vec(ha1, f_act, ha1); // ha1 BROKEN HERE !
      gsl_vector_add_constant(ha1, -1.0);
      gsl_vector_scale(ha1, -1.0);
      gsl_vector_mul(hd1, ha1);

    gsl_blas_dgemv(CblasTrans, 1.0, &view_dw2.matrix, d2, 1.0, hd1);
    gsl_blas_dgemv(CblasTrans, 1.0, w2, hd2, 1.0, hd1);
    gsl_vector_mul(hd1, hsig);
  
    // calculate dTH
    gsl_vector_set_zero( dTH );
    gsl_vector_memcpy( &view_dTHth1.vector, hd1 );
    gsl_vector_memcpy( &view_dTHth2.vector, hd2 );
    gsl_blas_dger( 1.0, hd1, x[n], &view_dTHw1.matrix );
    gsl_blas_dger( 1.0, hd2, z, &view_dTHw2.matrix );
    gsl_blas_dger( 1.0, d2, hz, &view_dTHw2.matrix );

    // add to sum
    gsl_vector_add(sum_dTH, dTH);
  }

  gsl_blas_ddot(d, sum_dTH, &numer);

  // calc eta
  eta = denom / numer;
}

bool Sd::learn( unsigned _times, double _eta, bool _reset )
{
  if(_reset){
    logErr.clear();
    logEta.clear();
  }

#define _MRT_INIT_
#ifdef _MRT_INIT_
  cout << "initialized by Mrt method" << endl;
  setWeightMrt();
#endif /* _MRT_INIT_ */

  cout << "w in the first state." << endl;
  show();

  double tmp;
  eta = _eta;
  logEta.push_back( eta );
  logErr.push_back( bbprop() );
  for(unsigned i = 0 ; i < _times; i++){

    // d(t) := -1.0 * g(t)
    gsl_vector_memcpy(d, g);
    gsl_vector_scale(d, -1.0);

    // calc eta(t)
    // renew w(t+1)
#if defined( _ARMIJO_ )
    // dw is abandoned
    armijo();
    logEta.push_back( eta );
    add_vec( w, alpha, dw );
#elif defined( _BPQ_ )
    bpq();
    logEta.push_back( eta );
    renewWeight();
#else /* Fixed Eta */
    logEta.push_back( eta );
    renewWeight();
#endif /* Choice of  */

    // calc g(t+1) and dg(t) (batch bprop)
    gsl_vector_memcpy(dg, g);
    tmp = bbprop();
    if( isnan(tmp) || isinf(tmp) ){
      cerr << "nan/inf appeared (" << i << ") by SD" << endl;
      return false;
    }
    logErr.push_back( tmp );
    gsl_vector_sub(dg, g);
    gsl_vector_scale(dg, -1.0);

    // save dada once every 1000 times
    if( i % 1000 == 0 ){
      saveError();
      saveEta();
    }
  }
  
  cout << "w in the last state" << endl;
  show();
  cout << "the last error is " << logErr[_times] << endl;
  return true;
}

bool Sd::saveEta()
{
  string fileError = projName + ".eta";
  ofstream fout( fileError.c_str() );
  if( !fout ){
    cerr << "Opening " << fileError << " failed." << endl;
    return false;
  }

  vector<double>::iterator e = logEta.begin();
  for(; e != logEta.end(); e++)
    fout << *e << endl;

  return true;
}

#ifdef DEBUG_SD
int main()
{
  cout << "DEBUG_SD dayo" << endl;
  Sd sd("data.dat", 4, "debug_sd", 0.0, 1.0);
  sd.learn( 100, 0.01 );
  sd.saveError();
  sd.saveTrial();
  sd.saveGraph();

  return 0;
}
#endif /* DEBUG_SD */
最終更新:2009年12月15日 18:44
ツールボックス

下から選んでください:

新しいページを作成する
ヘルプ / FAQ もご覧ください。