From 533562506ab781039e49d4d3fb27302e927b4873 Mon Sep 17 00:00:00 2001 From: KimSangYeon-DGU Date: Fri, 12 Apr 2019 05:15:51 +0900 Subject: [PATCH] Convert DiagonalGMMs to GMMs in mlpack_gmm_train --- src/mlpack/methods/gmm/gmm_train_main.cpp | 53 +++++++++++++++++++++-- 1 file changed, 49 insertions(+), 4 deletions(-) diff --git a/src/mlpack/methods/gmm/gmm_train_main.cpp b/src/mlpack/methods/gmm/gmm_train_main.cpp index a01adfcb0e..8bc64cddcf 100644 --- a/src/mlpack/methods/gmm/gmm_train_main.cpp +++ b/src/mlpack/methods/gmm/gmm_train_main.cpp @@ -14,6 +14,7 @@ #include #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 em(maxIterations, tolerance, k); - likelihood = gmm->Train(dataPoints, CLI::GetParam("trials"), false, + EMFit em(maxIterations, + tolerance, k); + + likelihood = dgmm.Train(dataPoints, CLI::GetParam("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, DiagonalConstraint> em(maxIterations, tolerance); - likelihood = gmm->Train(dataPoints, CLI::GetParam("trials"), false, + EMFit, PositiveDefiniteConstraint, + distribution::DiagonalGaussianDistribution> em(maxIterations, + tolerance); + + likelihood = dgmm.Train(dataPoints, CLI::GetParam("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) {