. 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:
+4
-4
@@ -184,14 +184,14 @@ arma::mat Rbm::sample(const arma::mat &src)
|
|||||||
|
|
||||||
arma::mat Rbm::toHidden(const arma::mat &visible)
|
arma::mat Rbm::toHidden(const arma::mat &visible)
|
||||||
{
|
{
|
||||||
arma::mat h = visible.t() * m_w + m_bh;
|
arma::mat h = visible * m_w + arma::repmat(m_bh, visible.n_rows, 1);
|
||||||
return h;
|
return probsLogistic(h);
|
||||||
}
|
}
|
||||||
|
|
||||||
arma::mat Rbm::toVisible(const arma::mat &hidden)
|
arma::mat Rbm::toVisible(const arma::mat &hidden)
|
||||||
{
|
{
|
||||||
arma::mat v = hidden * m_w.t() + m_bv;
|
arma::mat v = hidden * m_w.t() + arma::repmat(m_bv, hidden.n_rows, 1);
|
||||||
return v;
|
return probsLogistic(v);
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
|
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
|
||||||
|
|||||||
+1
-1
@@ -65,7 +65,7 @@ public:
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
Rbm(size_t numHidden, size_t numVisible);
|
Rbm(size_t numVisible, size_t numHidden);
|
||||||
Rbm(const Rbm& orig);
|
Rbm(const Rbm& orig);
|
||||||
virtual ~Rbm();
|
virtual ~Rbm();
|
||||||
|
|
||||||
|
|||||||
+4
-4
@@ -56,10 +56,10 @@ int main()
|
|||||||
Rbm::Params params;
|
Rbm::Params params;
|
||||||
Rbm rbm(28*28, 64);
|
Rbm rbm(28*28, 64);
|
||||||
arma::mat batch = arma::randu(100, 28*28);
|
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 v = arma::randu(10, 28*28);
|
||||||
// arma::mat h = rbm.toHidden(v);
|
arma::mat h = rbm.toHidden(v);
|
||||||
// arma::mat r = rbm.toVisible(h);
|
arma::mat r = rbm.toVisible(h);
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
Reference in New Issue
Block a user