- pass numContext

- fixed crash when numContext == 0

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@761 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-09 08:01:25 +00:00
parent 0ee0a732e9
commit c2ffecd66a
10 changed files with 46 additions and 21 deletions
+21 -3
View File
@@ -119,10 +119,13 @@ public:
virtual ~Rbm();
void weightsInit(double stddev, double mu=0.0);
void weightsAssign(const arma::mat &w)
void weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv)
{
m_whv.submat(0, 0, w.n_rows-1, w.n_cols-1) = w;
m_bhv.submat(0, 0, bhv.n_rows-1, bhv.n_cols-1) = bhv;
m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
}
void train(arma::mat const &batch, IListener *pListener=nullptr);
static arma::mat normalize(const arma::mat &hidden);
@@ -141,12 +144,27 @@ public:
arma::mat toHiddenProbs(const arma::mat &visible) const
{
return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx)));
return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx)));
}
arma::mat toVisibleProbs(const arma::mat &hidden) const
{
return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, m_bv.n_cols-m_bhv.n_cols);
return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, numVisible() - numContext());
}
size_t numContext() const
{
return m_ctx.size();
}
size_t numHidden() const
{
return m_bhv.size();
}
size_t numVisible() const
{
return m_bv.size();
}
private: