#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