Merge pull request #1860 from KimSangYeon-DGU/conversion

Convert DiagonalGMMs into GMMs in `mlpack_gmm_train`.
This commit is contained in:
Shikhar Jaiswal
2019-04-19 14:46:00 +05:30
committed by GitHub
+49 -4
View File
@@ -14,6 +14,7 @@
#include <mlpack/core/util/mlpack_main.hpp>
#include "gmm.hpp"
#include "diagonal_gmm.hpp"
#include "no_constraint.hpp"
#include "diagonal_constraint.hpp"
@@ -211,12 +212,34 @@ static void mlpackMain()
// to use different types.
if (diagonalCovariance)
{
// Convert GMMs into DiagonalGMMs.
DiagonalGMM dgmm(gmm->Gaussians(), gmm->Dimensionality());
for (size_t i = 0; i < size_t(gaussians); i++)
{
dgmm.Component(i).Mean() = gmm->Component(i).Mean();
dgmm.Component(i).Covariance(
std::move(arma::diagvec(gmm->Component(i).Covariance())));
}
dgmm.Weights() = gmm->Weights();
// Compute the parameters of the model using the EM algorithm.
Timer::Start("em");
EMFit<KMeansType, DiagonalConstraint> em(maxIterations, tolerance, k);
likelihood = gmm->Train(dataPoints, CLI::GetParam<int>("trials"), false,
EMFit<KMeansType, PositiveDefiniteConstraint,
distribution::DiagonalGaussianDistribution> em(maxIterations,
tolerance, k);
likelihood = dgmm.Train(dataPoints, CLI::GetParam<int>("trials"), false,
em);
Timer::Stop("em");
// Convert DiagonalGMMs into GMMs.
for (size_t i = 0; i < size_t(gaussians); i++)
{
gmm->Component(i).Mean() = dgmm.Component(i).Mean();
gmm->Component(i).Covariance(
std::move(arma::diagmat(dgmm.Component(i).Covariance())));
}
gmm->Weights() = dgmm.Weights();
}
else if (forcePositive)
{
@@ -243,12 +266,34 @@ static void mlpackMain()
// to use different types.
if (diagonalCovariance)
{
// Convert GMMs into DiagonalGMMs.
DiagonalGMM dgmm(gmm->Gaussians(), gmm->Dimensionality());
for (size_t i = 0; i < size_t(gaussians); i++)
{
dgmm.Component(i).Mean() = gmm->Component(i).Mean();
dgmm.Component(i).Covariance(
std::move(arma::diagvec(gmm->Component(i).Covariance())));
}
dgmm.Weights() = gmm->Weights();
// Compute the parameters of the model using the EM algorithm.
Timer::Start("em");
EMFit<kmeans::KMeans<>, DiagonalConstraint> em(maxIterations, tolerance);
likelihood = gmm->Train(dataPoints, CLI::GetParam<int>("trials"), false,
EMFit<kmeans::KMeans<>, PositiveDefiniteConstraint,
distribution::DiagonalGaussianDistribution> em(maxIterations,
tolerance);
likelihood = dgmm.Train(dataPoints, CLI::GetParam<int>("trials"), false,
em);
Timer::Stop("em");
// Convert DiagonalGMMs into GMMs.
for (size_t i = 0; i < size_t(gaussians); i++)
{
gmm->Component(i).Mean() = dgmm.Component(i).Mean();
gmm->Component(i).Covariance(
std::move(arma::diagmat(dgmm.Component(i).Covariance())));
}
gmm->Weights() = dgmm.Weights();
}
else if (forcePositive)
{