converted guassian to .hpp

This commit is contained in:
shubham1206agra
2022-04-29 12:10:55 +05:30
parent 3d18d41a5e
commit 2a451e3047
3 changed files with 26 additions and 22 deletions
+1 -1
View File
@@ -4,7 +4,7 @@ set(SOURCES
discrete_distribution.hpp
discrete_distribution_impl.hpp
gaussian_distribution.hpp
gaussian_distribution.cpp
gaussian_distribution_impl.hpp
laplace_distribution.hpp
laplace_distribution.cpp
regression_distribution.hpp
@@ -16,7 +16,7 @@
#include <mlpack/prereqs.hpp>
namespace mlpack {
namespace distribution {
namespace distribution /** Probability distributions. */ {
/**
* A single multivariate Gaussian distribution.
@@ -194,4 +194,7 @@ class GaussianDistribution
} // namespace distribution
} // namespace mlpack
// Include implementation.
#include "gaussian_distribution_impl.hpp"
#endif
@@ -1,5 +1,5 @@
/**
* @file core/dists/gaussian_distribution.cpp
* @file core/dists/gaussian_distribution_impl.hpp
* @author Ryan Curtin
* @author Michael Fox
*
@@ -10,33 +10,35 @@
* 3-clause BSD license along with mlpack. If not, see
* http://www.opensource.org/licenses/BSD-3-Clause for more information.
*/
#ifndef MLPACK_CORE_DISTRIBUTIONS_GAUSSIAN_DISTRIBUTION_IMPL_HPP
#define MLPACK_CORE_DISTRIBUTIONS_GAUSSIAN_DISTRIBUTION_IMPL_HPP
#include "gaussian_distribution.hpp"
#include <mlpack/methods/gmm/positive_definite_constraint.hpp>
using namespace mlpack;
using namespace mlpack::distribution;
namespace mlpack {
namespace distribution /** Probability distributions. */ {
GaussianDistribution::GaussianDistribution(const arma::vec& mean,
inline GaussianDistribution::GaussianDistribution(const arma::vec& mean,
const arma::mat& covariance)
: mean(mean), logDetCov(0.0)
{
Covariance(covariance);
}
void GaussianDistribution::Covariance(const arma::mat& covariance)
inline void GaussianDistribution::Covariance(const arma::mat& covariance)
{
this->covariance = covariance;
FactorCovariance();
}
void GaussianDistribution::Covariance(arma::mat&& covariance)
inline void GaussianDistribution::Covariance(arma::mat&& covariance)
{
this->covariance = std::move(covariance);
FactorCovariance();
}
void GaussianDistribution::FactorCovariance()
inline void GaussianDistribution::FactorCovariance()
{
// On Armadillo < 4.500, the "lower" option isn't available.
@@ -68,7 +70,7 @@ void GaussianDistribution::FactorCovariance()
logDetCov *= 2;
}
double GaussianDistribution::LogProbability(const arma::vec& observation) const
inline double GaussianDistribution::LogProbability(const arma::vec& observation) const
{
const size_t k = observation.n_elem;
const arma::vec diff = mean - observation;
@@ -76,7 +78,7 @@ double GaussianDistribution::LogProbability(const arma::vec& observation) const
return -0.5 * k * log2pi - 0.5 * logDetCov - 0.5 * v(0);
}
arma::vec GaussianDistribution::Random() const
inline arma::vec GaussianDistribution::Random() const
{
return covLower * arma::randn<arma::vec>(mean.n_elem) + mean;
}
@@ -86,7 +88,7 @@ arma::vec GaussianDistribution::Random() const
*
* @param observations List of observations.
*/
void GaussianDistribution::Train(const arma::mat& observations)
inline void GaussianDistribution::Train(const arma::mat& observations)
{
if (observations.n_cols > 0)
{
@@ -95,10 +97,7 @@ void GaussianDistribution::Train(const arma::mat& observations)
}
else // This will end up just being empty.
{
// TODO(stephentu): why do we allow this case? why not throw an error?
mean.zeros(0);
covariance.zeros(0);
return;
Log::Fatal << "Observation columns equal to 0." << std::endl;
}
// Calculate the mean.
@@ -130,7 +129,7 @@ void GaussianDistribution::Train(const arma::mat& observations)
* account the probability of each observation actually being from this
* distribution.
*/
void GaussianDistribution::Train(const arma::mat& observations,
inline void GaussianDistribution::Train(const arma::mat& observations,
const arma::vec& probabilities)
{
if (observations.n_cols > 0)
@@ -140,10 +139,7 @@ void GaussianDistribution::Train(const arma::mat& observations,
}
else // This will end up just being empty.
{
// TODO(stephentu): same as above
mean.zeros(0);
covariance.zeros(0);
return;
Log::Fatal << "Observation columns equal to 0." << std::endl;
}
double sumProb = 0;
@@ -185,3 +181,8 @@ void GaussianDistribution::Train(const arma::mat& observations,
FactorCovariance();
}
} // namespace distribution
} // namespace mlpack
#endif