. improved toHidden(), toVisible()

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@564 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-22 06:20:43 +00:00
parent 8d323ee054
commit 675a5485fe
3 changed files with 9 additions and 9 deletions
+4 -4
View File
@@ -184,14 +184,14 @@ arma::mat Rbm::sample(const arma::mat &src)
arma::mat Rbm::toHidden(const arma::mat &visible)
{
arma::mat h = visible.t() * m_w + m_bh;
return h;
arma::mat h = visible * m_w + arma::repmat(m_bh, visible.n_rows, 1);
return probsLogistic(h);
}
arma::mat Rbm::toVisible(const arma::mat &hidden)
{
arma::mat v = hidden * m_w.t() + m_bv;
return v;
arma::mat v = hidden * m_w.t() + arma::repmat(m_bv, hidden.n_rows, 1);
return probsLogistic(v);
}
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
+1 -1
View File
@@ -65,7 +65,7 @@ public:
}
};
Rbm(size_t numHidden, size_t numVisible);
Rbm(size_t numVisible, size_t numHidden);
Rbm(const Rbm& orig);
virtual ~Rbm();
+4 -4
View File
@@ -56,10 +56,10 @@ int main()
Rbm::Params params;
Rbm rbm(28*28, 64);
arma::mat batch = arma::randu(100, 28*28);
rbm.train(batch, 100, 10, params, &statusDisplay);
rbm.train(batch, 10, 10, params, &statusDisplay);
// arma::mat v = arma::randu(28*28, 1);
// arma::mat h = rbm.toHidden(v);
// arma::mat r = rbm.toVisible(h);
arma::mat v = arma::randu(10, 28*28);
arma::mat h = rbm.toHidden(v);
arma::mat r = rbm.toVisible(h);
return 0;
}