- 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) , m_bh(bh)
{ {
Noise_Init(&m_noise, 0x32727155); Noise_Init(&m_noise, 0x32727155);
uniform(m_w, 0.0, m_params.m_weightInit); uniform(m_w, 0.0, m_params.weightInit);
uniform(m_bh, 0.0, m_params.m_weightInit); uniform(m_bh, 0.0, m_params.weightInit);
uniform(m_bv, 0.0, m_params.m_weightInit); uniform(m_bv, 0.0, m_params.weightInit);
} }
Rbm::Rbm(const Rbm& orig) 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); arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
trainingSizeRemain -= miniBatchSizeActual; trainingSizeRemain -= miniBatchSizeActual;
batchRowIndex += miniBatchSizeActual; batchRowIndex += miniBatchSizeActual;
double learning_rate = m_params.m_learningRate/std::min(miniBatchSizeActual, trainingSize); double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize);
double weight_decay = m_params.m_weightDecay/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_state(miniBatchSizeActual, m_w.n_rows);
arma::mat vis_probs(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 // 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 // When the hidden units are being driven by data, always use stochastic binary states
vis_state = sample(miniBatch); 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); hid_probs = probsLogistic(hid_state);
// Sample hidden // Sample hidden
if (m_params.m_doRaoBlackwell) if (m_params.doRaoBlackwell)
{ {
hid_state = hid_probs; 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_v = sum(vis_state, 0);
grad_bias_h = sum(hid_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 // Create visible reconstruction (a fantasy...) given hid
if (m_params.m_gibbsDoSampleHidden) if (m_params.gibbsDoSampleHidden)
{ {
vis_probs = toVisibleProbs(sample(hid_probs)); 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 // Create hidden representation given v
if (m_params.m_gibbsDoSampleVisible) if (m_params.gibbsDoSampleVisible)
{ {
hid_state = toHiddenState(sample(vis_probs)); 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.L1 = accu(abs(m_w));
status.L2 = accu(m_w % m_w); status.L2 = accu(m_w % m_w);
momentum_bias_v = m_params.m_momentum*momentum_bias_v + grad_bias_v; momentum_bias_v = m_params.momentum*momentum_bias_v + grad_bias_v;
momentum_bias_h = m_params.m_momentum*momentum_bias_h + grad_bias_h; momentum_bias_h = m_params.momentum*momentum_bias_h + grad_bias_h;
momentum_weights = m_params.m_momentum*momentum_weights + grad_weight - status.L2*penalty_weights; momentum_weights = m_params.momentum*momentum_weights + grad_weight - status.L2*penalty_weights;
m_bv += learning_rate*momentum_bias_v; m_bv += learning_rate*momentum_bias_v;
m_bh += learning_rate*momentum_bias_h; m_bh += learning_rate*momentum_bias_h;
+18 -18
View File
@@ -24,27 +24,27 @@ public:
struct Params struct Params
{ {
Params() Params()
: m_weightInit(0.01) : weightInit(0.01)
, m_weightDecay(0.001) , weightDecay(0.001)
, m_learningRate(0.1) , learningRate(0.1)
, m_momentum(0.5) , momentum(0.5)
, m_doRaoBlackwell(true) , doRaoBlackwell(true)
, m_gibbsDoSampleVisible(false) , gibbsDoSampleVisible(false)
, m_gibbsDoSampleHidden(true) , gibbsDoSampleHidden(true)
, m_doSampleBatch(false) , doSampleBatch(false)
, m_numGibbs(1) , numGibbs(1)
{ {
} }
double m_weightInit; double weightInit;
double m_weightDecay; double weightDecay;
double m_learningRate; double learningRate;
double m_momentum; double momentum;
bool m_doRaoBlackwell; bool doRaoBlackwell;
bool m_gibbsDoSampleVisible; bool gibbsDoSampleVisible;
bool m_gibbsDoSampleHidden; bool gibbsDoSampleHidden;
bool m_doSampleBatch; bool doSampleBatch;
size_t m_numGibbs; size_t numGibbs;
}; };
struct Status struct Status