Fix Tests and Clean Implementation
This commit is contained in:
@@ -111,8 +111,11 @@ typename std::enable_if<std::is_same<Policy, BinaryRBM>::value, double>::type
|
||||
RBM<InitializationRuleType, DataType, PolicyType>::FreeEnergy(
|
||||
arma::Mat<ElemType>&& input)
|
||||
{
|
||||
preActivation = arma::log(1 + arma::trunc_exp((weight * input) + hiddenBias));
|
||||
return -(arma::accu(preActivation) + arma::dot(input, visibleBias));
|
||||
preActivation = (weight * input);
|
||||
preActivation.each_col() += hiddenBias;
|
||||
preActivation = arma::log(1 + arma::trunc_exp(preActivation));
|
||||
return -(arma::accu(preActivation) + arma::dot(input, arma::repmat(
|
||||
visibleBias, 1, input.n_cols)));
|
||||
}
|
||||
|
||||
template<
|
||||
@@ -204,7 +207,8 @@ typename std::enable_if<std::is_same<Policy, BinaryRBM>::value, void>::type
|
||||
RBM<InitializationRuleType, DataType, PolicyType>::VisibleMean(DataType&& input,
|
||||
DataType&& output)
|
||||
{
|
||||
output = weight.t() * input + visibleBias;
|
||||
output = weight.t() * input;
|
||||
output.each_col() += visibleBias;
|
||||
LogisticFunction::Fn(output, output);
|
||||
}
|
||||
|
||||
@@ -218,7 +222,8 @@ typename std::enable_if<std::is_same<Policy, BinaryRBM>::value, void>::type
|
||||
RBM<InitializationRuleType, DataType, PolicyType>::HiddenMean(DataType&& input,
|
||||
DataType&& output)
|
||||
{
|
||||
output = weight * input + hiddenBias;
|
||||
output = weight * input;
|
||||
output.each_col() += hiddenBias;
|
||||
LogisticFunction::Fn(output, output);
|
||||
}
|
||||
|
||||
@@ -281,6 +286,7 @@ void RBM<InitializationRuleType, DataType, PolicyType>::Gradient(
|
||||
|
||||
negativeGradient += tempNegativeGradient;
|
||||
}
|
||||
|
||||
gradient = ((negativeGradient / negSteps) - positiveGradient);
|
||||
}
|
||||
|
||||
@@ -313,6 +319,8 @@ void RBM<InitializationRuleType, DataType, PolicyType>::serialize(
|
||||
ar & BOOST_SERIALIZATION_NVP(negSteps);
|
||||
ar & BOOST_SERIALIZATION_NVP(persistence);
|
||||
ar & BOOST_SERIALIZATION_NVP(poolSize);
|
||||
ar & BOOST_SERIALIZATION_NVP(visibleBias);
|
||||
ar & BOOST_SERIALIZATION_NVP(hiddenBias);
|
||||
ar & BOOST_SERIALIZATION_NVP(weight);
|
||||
ar & BOOST_SERIALIZATION_NVP(weightCube);
|
||||
ar & BOOST_SERIALIZATION_NVP(spikeBias);
|
||||
|
||||
@@ -41,9 +41,8 @@ RBM<InitializationRuleType, DataType, PolicyType>::Reset()
|
||||
slabMean.set_size(poolSize, hiddenSize);
|
||||
|
||||
// Weight shape D * K * N
|
||||
weightCube = arma::Cube<ElemType>(parameter.memptr(),
|
||||
visibleSize, poolSize, hiddenSize,
|
||||
false, false);
|
||||
weightCube = arma::Cube<ElemType>(parameter.memptr(), visibleSize, poolSize,
|
||||
hiddenSize, false, false);
|
||||
// spike bias shape N * 1
|
||||
spikeBias = DataType(parameter.memptr() + weight.n_elem, hiddenSize, 1,
|
||||
false, false);
|
||||
@@ -70,8 +69,7 @@ typename std::enable_if<std::is_same<Policy, SpikeSlabRBM>::value, double>::type
|
||||
RBM<InitializationRuleType, DataType, PolicyType>::FreeEnergy(
|
||||
arma::Mat<ElemType>&& input)
|
||||
{
|
||||
ElemType freeEnergy = 0.5 * arma::as_scalar(visiblePenalty(0) * input.t() *
|
||||
input);
|
||||
ElemType freeEnergy = 0.5 * visiblePenalty(0) * arma::accu(input.t() * input);
|
||||
|
||||
freeEnergy -= 0.5 * hiddenSize * poolSize *
|
||||
std::log((2.0 * M_PI) / slabPenalty);
|
||||
@@ -112,11 +110,15 @@ RBM<InitializationRuleType, DataType, PolicyType>::Phase(
|
||||
SlabMean(std::move(input), std::move(spikeSamples), std::move(slabMean));
|
||||
|
||||
for (size_t i = 0 ; i < hiddenSize; i++)
|
||||
weightGrad.slice(i) = input * slabMean.col(i).t() * spikeMean(i);
|
||||
{
|
||||
weightGrad.slice(i) = input * arma::repmat(slabMean.col(i).t(),
|
||||
input.n_cols, 1) * spikeMean(i);
|
||||
}
|
||||
|
||||
spikeBiasGrad = spikeMean;
|
||||
|
||||
visiblePenaltyGrad = -0.5 * input.t() * input;
|
||||
visiblePenaltyGrad = -0.5 * arma::accu(input.t() * input)
|
||||
/ std::pow(input.n_cols, 2);
|
||||
}
|
||||
|
||||
template<
|
||||
@@ -166,7 +168,9 @@ RBM<InitializationRuleType, DataType, PolicyType>::SampleVisible(
|
||||
output(i) = math::RandNormal(visibleMean(i), 1.0 / visiblePenalty(0));
|
||||
}
|
||||
if (arma::norm(output, 2) < radius)
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (k == numMaxTrials)
|
||||
@@ -197,7 +201,9 @@ RBM<InitializationRuleType, DataType, PolicyType>::VisibleMean(
|
||||
false);
|
||||
|
||||
for (size_t i = 0; i < hiddenSize; i++)
|
||||
{
|
||||
output += weightCube.slice(i) * slab.col(i) * spike(i);
|
||||
}
|
||||
|
||||
output = ((1.0 / visiblePenalty(0)) * output);
|
||||
}
|
||||
@@ -237,9 +243,9 @@ RBM<InitializationRuleType, DataType, PolicyType>::SpikeMean(
|
||||
{
|
||||
for (size_t i = 0; i < hiddenSize; i++)
|
||||
{
|
||||
spikeMean(i) = LogisticFunction::Fn(0.5 * (1.0 / slabPenalty) *
|
||||
arma::as_scalar(visible.t() * weightCube.slice(i) *
|
||||
weightCube.slice(i).t() * visible) + spikeBias(i));
|
||||
spikeMean(i) = LogisticFunction::Fn(0.5 * (1.0 / slabPenalty) * arma::accu(
|
||||
visible.t() * (weightCube.slice(i) * weightCube.slice(i).t()) * visible)
|
||||
/ std::pow(visible.n_cols, 2) + spikeBias(i));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -255,7 +261,9 @@ RBM<InitializationRuleType, DataType, PolicyType>::SampleSpike(
|
||||
DataType&& spike)
|
||||
{
|
||||
for (size_t i = 0; i < hiddenSize; i++)
|
||||
{
|
||||
spike(i) = math::RandBernoulli(spikeMean(i));
|
||||
}
|
||||
}
|
||||
|
||||
template<
|
||||
@@ -272,8 +280,8 @@ RBM<InitializationRuleType, DataType, PolicyType>::SlabMean(
|
||||
{
|
||||
for (size_t i = 0; i < hiddenSize; i++)
|
||||
{
|
||||
slabMean.col(i) = (1.0 / slabPenalty) * spike(i) *
|
||||
weightCube.slice(i).t() * visible;
|
||||
slabMean.col(i) = arma::mean((1.0 / slabPenalty) * spike(i) *
|
||||
weightCube.slice(i).t() * visible, 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -296,6 +304,7 @@ RBM<InitializationRuleType, DataType, PolicyType>::SampleSlab(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace ann
|
||||
} // namespace mlpack
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ using namespace mlpack::regression;
|
||||
|
||||
BOOST_AUTO_TEST_SUITE(RBMNetworkTest);
|
||||
|
||||
BOOST_AUTO_TEST_CASE(ClassificationTest)
|
||||
BOOST_AUTO_TEST_CASE(BinaryRBMClassificationTest)
|
||||
{
|
||||
// Normalised dataset.
|
||||
int hiddenLayerSize = 100;
|
||||
@@ -74,7 +74,7 @@ BOOST_AUTO_TEST_CASE(ClassificationTest)
|
||||
|
||||
size_t numRBMIterations = trainData.n_cols * numEpoches;
|
||||
numRBMIterations /= batchSize;
|
||||
optimization::StandardSGD msgd(0.06, batchSize, numRBMIterations, 0, true);
|
||||
optimization::StandardSGD msgd(0.03, batchSize, numRBMIterations, 0, true);
|
||||
model.Reset();
|
||||
model.VisibleBias().ones();
|
||||
model.HiddenBias().ones();
|
||||
|
||||
Reference in New Issue
Block a user