[RBM]
- miniBatchSize is parameter of train() git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@307 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+9
-15
@@ -17,7 +17,6 @@ Rbm::Rbm(Weights &weights, const MatrixXd &batch)
|
||||
, m_variableSigma(weights.getNumVisible())
|
||||
, m_progress(0)
|
||||
{
|
||||
setMiniBatchSize(batch.rows());
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
m_variableSigma.fill(m_params.m_constantSigma);
|
||||
updateHiddenBatch();
|
||||
@@ -198,17 +197,17 @@ MatrixXd Rbm::calcZ(MatrixXd &v, MatrixXd &h)
|
||||
return t1;
|
||||
}
|
||||
|
||||
void Rbm::train(uint32_t numEpochs, double sigmaMin)
|
||||
void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
{
|
||||
uint32_t i;
|
||||
uint32_t epoch;
|
||||
uint32_t gibbs;
|
||||
size_t i;
|
||||
size_t epoch;
|
||||
size_t gibbs;
|
||||
|
||||
size_t trainingSize = m_batch.rows();
|
||||
size_t trainingSizeRemain = trainingSize;
|
||||
size_t batchRowIndex = 0;
|
||||
|
||||
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(m_params.m_miniBatchSize, trainingSize));
|
||||
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
|
||||
|
||||
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
@@ -221,14 +220,14 @@ void Rbm::train(uint32_t numEpochs, double sigmaMin)
|
||||
while (trainingSizeRemain)
|
||||
{
|
||||
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
|
||||
size_t toSlice = std::min(m_params.m_miniBatchSize, trainingSizeRemain);
|
||||
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
||||
MatrixXd batch = m_batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible());
|
||||
trainingSizeRemain -= toSlice;
|
||||
batchRowIndex += toSlice;
|
||||
size_t batchSize = batch.rows();
|
||||
double mu_w = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
||||
double mu_biasV = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
||||
double mu_biasH = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
||||
double mu_w = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||
double mu_biasV = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||
double mu_biasH = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||
|
||||
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
||||
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
||||
@@ -514,11 +513,6 @@ void Rbm::setNumGibbs(size_t value)
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void Rbm::setMiniBatchSize(size_t size)
|
||||
{
|
||||
m_params.m_miniBatchSize = size;
|
||||
}
|
||||
|
||||
void Rbm::setMuWeights(double value)
|
||||
{
|
||||
m_params.m_muWeights = value;
|
||||
|
||||
Reference in New Issue
Block a user