- 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)
|
, 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
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user