- renamed Params member
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@570 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+13
-13
@@ -21,9 +21,9 @@ Rbm::Rbm(const Params& params, arma::mat &w, arma::mat &bv, arma::mat &bh)
|
||||
, m_bh(bh)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
uniform(m_w, 0.0, m_params.m_weightInit);
|
||||
uniform(m_bh, 0.0, m_params.m_weightInit);
|
||||
uniform(m_bv, 0.0, m_params.m_weightInit);
|
||||
uniform(m_w, 0.0, m_params.weightInit);
|
||||
uniform(m_bh, 0.0, m_params.weightInit);
|
||||
uniform(m_bv, 0.0, m_params.weightInit);
|
||||
}
|
||||
|
||||
Rbm::Rbm(const Rbm& orig)
|
||||
@@ -67,8 +67,8 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||
trainingSizeRemain -= miniBatchSizeActual;
|
||||
batchRowIndex += miniBatchSizeActual;
|
||||
double learning_rate = m_params.m_learningRate/std::min(miniBatchSizeActual, trainingSize);
|
||||
double weight_decay = m_params.m_weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
||||
double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize);
|
||||
double weight_decay = m_params.weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
||||
|
||||
arma::mat vis_state(miniBatchSizeActual, m_w.n_rows);
|
||||
arma::mat vis_probs(miniBatchSizeActual, m_w.n_rows);
|
||||
@@ -79,7 +79,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
{
|
||||
|
||||
// Create hidden layer base on training data
|
||||
if (m_params.m_doSampleBatch)
|
||||
if (m_params.doSampleBatch)
|
||||
{
|
||||
// When the hidden units are being driven by data, always use stochastic binary states
|
||||
vis_state = sample(miniBatch);
|
||||
@@ -93,7 +93,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
hid_probs = probsLogistic(hid_state);
|
||||
|
||||
// Sample hidden
|
||||
if (m_params.m_doRaoBlackwell)
|
||||
if (m_params.doRaoBlackwell)
|
||||
{
|
||||
hid_state = hid_probs;
|
||||
}
|
||||
@@ -107,10 +107,10 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
grad_bias_v = sum(vis_state, 0);
|
||||
grad_bias_h = sum(hid_state, 0);
|
||||
|
||||
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
|
||||
for (gibbs=0; gibbs < m_params.numGibbs; gibbs++)
|
||||
{
|
||||
// Create visible reconstruction (a fantasy...) given hid
|
||||
if (m_params.m_gibbsDoSampleHidden)
|
||||
if (m_params.gibbsDoSampleHidden)
|
||||
{
|
||||
vis_probs = toVisibleProbs(sample(hid_probs));
|
||||
}
|
||||
@@ -120,7 +120,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
}
|
||||
|
||||
// Create hidden representation given v
|
||||
if (m_params.m_gibbsDoSampleVisible)
|
||||
if (m_params.gibbsDoSampleVisible)
|
||||
{
|
||||
hid_state = toHiddenState(sample(vis_probs));
|
||||
}
|
||||
@@ -141,9 +141,9 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
|
||||
status.L1 = accu(abs(m_w));
|
||||
status.L2 = accu(m_w % m_w);
|
||||
momentum_bias_v = m_params.m_momentum*momentum_bias_v + grad_bias_v;
|
||||
momentum_bias_h = m_params.m_momentum*momentum_bias_h + grad_bias_h;
|
||||
momentum_weights = m_params.m_momentum*momentum_weights + grad_weight - status.L2*penalty_weights;
|
||||
momentum_bias_v = m_params.momentum*momentum_bias_v + grad_bias_v;
|
||||
momentum_bias_h = m_params.momentum*momentum_bias_h + grad_bias_h;
|
||||
momentum_weights = m_params.momentum*momentum_weights + grad_weight - status.L2*penalty_weights;
|
||||
|
||||
m_bv += learning_rate*momentum_bias_v;
|
||||
m_bh += learning_rate*momentum_bias_h;
|
||||
|
||||
+18
-18
@@ -24,27 +24,27 @@ public:
|
||||
struct Params
|
||||
{
|
||||
Params()
|
||||
: m_weightInit(0.01)
|
||||
, m_weightDecay(0.001)
|
||||
, m_learningRate(0.1)
|
||||
, m_momentum(0.5)
|
||||
, m_doRaoBlackwell(true)
|
||||
, m_gibbsDoSampleVisible(false)
|
||||
, m_gibbsDoSampleHidden(true)
|
||||
, m_doSampleBatch(false)
|
||||
, m_numGibbs(1)
|
||||
: weightInit(0.01)
|
||||
, weightDecay(0.001)
|
||||
, learningRate(0.1)
|
||||
, momentum(0.5)
|
||||
, doRaoBlackwell(true)
|
||||
, gibbsDoSampleVisible(false)
|
||||
, gibbsDoSampleHidden(true)
|
||||
, doSampleBatch(false)
|
||||
, numGibbs(1)
|
||||
{
|
||||
}
|
||||
|
||||
double m_weightInit;
|
||||
double m_weightDecay;
|
||||
double m_learningRate;
|
||||
double m_momentum;
|
||||
bool m_doRaoBlackwell;
|
||||
bool m_gibbsDoSampleVisible;
|
||||
bool m_gibbsDoSampleHidden;
|
||||
bool m_doSampleBatch;
|
||||
size_t m_numGibbs;
|
||||
double weightInit;
|
||||
double weightDecay;
|
||||
double learningRate;
|
||||
double momentum;
|
||||
bool doRaoBlackwell;
|
||||
bool gibbsDoSampleVisible;
|
||||
bool gibbsDoSampleHidden;
|
||||
bool doSampleBatch;
|
||||
size_t numGibbs;
|
||||
};
|
||||
|
||||
struct Status
|
||||
|
||||
Reference in New Issue
Block a user