- simplified penalty_weights
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@563 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+1
-21
@@ -132,20 +132,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
grad_bias_h -= sum(hid_probs, 0);
|
||||
grad_weight -= vis_probs.t() * hid_probs;
|
||||
|
||||
for (int i=0; i < m_w.n_rows; i++)
|
||||
{
|
||||
for (int j=0; j < m_w.n_cols; j++)
|
||||
{
|
||||
if (m_w(i,j) >= 0)
|
||||
{
|
||||
penalty_weights(i,j) = weight_decay;
|
||||
}
|
||||
else
|
||||
{
|
||||
penalty_weights(i,j) = -weight_decay;
|
||||
}
|
||||
}
|
||||
}
|
||||
penalty_weights = weight_decay*arma::sign(m_w);
|
||||
|
||||
status.L1 = accu(abs(m_w));
|
||||
status.L2 = accu(m_w % m_w);
|
||||
@@ -207,13 +194,6 @@ arma::mat Rbm::toVisible(const arma::mat &hidden)
|
||||
return v;
|
||||
}
|
||||
|
||||
arma::mat Rbm::uniform(size_t numRows, size_t numCols, double mu, double stdDev)
|
||||
{
|
||||
arma::mat dst = arma::zeros<arma::mat>(numRows, numCols);
|
||||
uniform(dst);
|
||||
return dst;
|
||||
}
|
||||
|
||||
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
|
||||
{
|
||||
for (size_t i=0; i < srcDst.n_rows; i++)
|
||||
|
||||
Reference in New Issue
Block a user