- 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:
2019-10-24 18:51:45 +00:00
parent f7acfc5bd0
commit d6ed27f5d2
2 changed files with 31 additions and 31 deletions
+13 -13
View File
@@ -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
View File
@@ -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