- add Rbm::normalize()
- normalize of training data git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@660 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -100,7 +100,7 @@ private:
|
||||
|
||||
void loadTraining()
|
||||
{
|
||||
m_stack->loadTraining();
|
||||
m_stack->loadTraining(rbmNormalizeDataToggleButton->getToggleState());
|
||||
patterSlider->setRange(0, m_stack->numTraining()-1, 1);
|
||||
}
|
||||
|
||||
|
||||
@@ -243,6 +243,18 @@ arma::mat Rbm::toVisibleProbs(const arma::mat &hidden) const
|
||||
return probsLogistic(toVisibleState(hidden));
|
||||
}
|
||||
|
||||
arma::mat Rbm::normalize(const arma::mat& src)
|
||||
{
|
||||
double mean = arma::accu(src)/src.n_elem;
|
||||
arma::mat x = src - mean;
|
||||
arma::mat x2 = x % x;
|
||||
double stddev = sqrt(arma::accu(x2)/x2.n_elem);
|
||||
|
||||
std::cout << "mean" << " : " << std::endl << mean << std::endl;
|
||||
std::cout << "stddev" << ": " << std::endl << stddev << std::endl;
|
||||
return x/stddev;
|
||||
}
|
||||
|
||||
void Rbm::uniform(arma::mat& srcDst, double stdDev, double mu)
|
||||
{
|
||||
#if 1
|
||||
|
||||
@@ -125,6 +125,7 @@ public:
|
||||
arma::mat toVisibleState(const arma::mat &hidden) const;
|
||||
arma::mat toHiddenProbs(const arma::mat &visible) const;
|
||||
arma::mat toVisibleProbs(const arma::mat &hidden) const;
|
||||
static arma::mat normalize(const arma::mat &hidden);
|
||||
const arma::mat& w() const;
|
||||
const arma::mat& bv() const;
|
||||
const arma::mat& bh() const;
|
||||
|
||||
+6
-1
@@ -222,7 +222,7 @@ arma::mat Stack::trainingData(Layer* pThatLayer)
|
||||
return thisBatch;
|
||||
}
|
||||
|
||||
size_t Stack::loadTraining()
|
||||
size_t Stack::loadTraining(bool doNormalize)
|
||||
{
|
||||
uint32_t numTraining = 0;
|
||||
uint32_t numVisible = 0;
|
||||
@@ -263,6 +263,11 @@ size_t Stack::loadTraining()
|
||||
}
|
||||
fclose(pFile);
|
||||
std::cout << "Loaded " << numTraining << " training samples\n";
|
||||
|
||||
if (doNormalize)
|
||||
{
|
||||
m_trainingData = Rbm::normalize(m_trainingData);
|
||||
}
|
||||
return numTraining;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -54,7 +54,7 @@ public:
|
||||
size_t numTraining();
|
||||
void addTraining(const arma::mat &toAdd);
|
||||
void delTraining(int index);
|
||||
size_t loadTraining();
|
||||
size_t loadTraining(bool doNormalize=false);
|
||||
size_t saveTraining();
|
||||
arma::mat& trainingData();
|
||||
arma::mat trainingData(Layer *pLayer);
|
||||
|
||||
Reference in New Issue
Block a user