Add --random_initialization for mlpack_hmm_train.

This commit is contained in:
Ryan Curtin
2016-05-16 15:11:56 -04:00
parent 274462840d
commit 3e4f3cade8
2 changed files with 66 additions and 1 deletions
+3
View File
@@ -22,6 +22,9 @@
mlpack_allknn and mlpack_allkfn programs will remain as copies until mlpack
3.0.0.
* Add --random_initialization option to mlpack_hmm_train, for use when no
labels are provided.
### mlpack 2.0.1
###### 2016-02-04
* Fix CMake to properly detect when MKL is being used with Armadillo.
+63 -1
View File
@@ -24,7 +24,9 @@ PROGRAM_INFO("Hidden Markov Model (HMM) Training", "This program allows a "
"\n\n"
"The HMM is trained with the Baum-Welch algorithm if no labels are "
"provided. The tolerance of the Baum-Welch algorithm can be set with the "
"--tolerance option."
"--tolerance option. In general it is a good idea to use random "
"initialization in this case, which can be specified with the "
"--random_initialization (-r) option."
"\n\n"
"Optionally, a pre-created HMM model can be used as a guess for the "
"transition matrix and emission probabilities; this is specifiable with "
@@ -47,6 +49,8 @@ PARAM_STRING("output_model_file", "File to save trained HMM to.", "o",
"output_hmm.xml");
PARAM_INT("seed", "Random seed. If 0, 'std::time(NULL)' is used.", "s", 0);
PARAM_DOUBLE("tolerance", "Tolerance of the Baum-Welch algorithm.", "T", 1e-5);
PARAM_FLAG("random_initialization", "Initialize emissions and transition "
"matrices with a uniform random distribution.", "r");
using namespace mlpack;
using namespace mlpack::hmm;
@@ -296,6 +300,21 @@ int main(int argc, char** argv)
HMM<DiscreteDistribution> hmm(size_t(states),
DiscreteDistribution(maxEmission), tolerance);
// Initialize with random starting point.
if (CLI::HasParam("random_initialization"))
{
hmm.Transition().randu();
for (size_t c = 0; c < hmm.Transition().n_cols; ++c)
hmm.Transition().col(c) /= arma::accu(hmm.Transition().col(c));
for (size_t e = 0; e < hmm.Emission().size(); ++e)
{
hmm.Emission()[e].Probabilities().randu();
hmm.Emission()[e].Probabilities() /=
arma::accu(hmm.Emission()[e].Probabilities());
}
}
// Now train it. Pass the already-loaded training data.
Train::Apply(hmm, &trainSeq);
}
@@ -314,6 +333,22 @@ int main(int argc, char** argv)
HMM<GaussianDistribution> hmm(size_t(states),
GaussianDistribution(dimensionality), tolerance);
// Initialize with random starting point.
if (CLI::HasParam("random_initialization"))
{
hmm.Transition().randu();
for (size_t c = 0; c < hmm.Transition().n_cols; ++c)
hmm.Transition().col(c) /= arma::accu(hmm.Transition().col(c));
for (size_t e = 0; e < hmm.Emission().size(); ++e)
{
hmm.Emission()[e].Mean().randu();
// Generate random covariance.
arma::mat r = arma::randu<arma::mat>(dimensionality, dimensionality);
hmm.Emission()[e].Covariance(r * r.t());
}
}
// Now train it.
Train::Apply(hmm, &trainSeq);
}
@@ -336,6 +371,33 @@ int main(int argc, char** argv)
HMM<GMM> hmm(size_t(states), GMM(size_t(gaussians), dimensionality),
tolerance);
// Initialize with random starting point.
if (CLI::HasParam("random_initialization"))
{
hmm.Transition().randu();
for (size_t c = 0; c < hmm.Transition().n_cols; ++c)
hmm.Transition().col(c) /= arma::accu(hmm.Transition().col(c));
for (size_t e = 0; e < hmm.Emission().size(); ++e)
{
// Random weights.
hmm.Emission()[e].Weights().randu();
hmm.Emission()[e].Weights() /=
arma::accu(hmm.Emission()[e].Weights());
// Random means and covariances.
for (int g = 0; g < gaussians; ++g)
{
hmm.Emission()[e].Component(g).Mean().randu();
// Generate random covariance.
arma::mat r = arma::randu<arma::mat>(dimensionality,
dimensionality);
hmm.Emission()[e].Component(g).Covariance(r * r.t());
}
}
}
// Issue a warning if the user didn't give labels.
if (!CLI::HasParam("labels_file"))
Log::Warn << "Unlabeled training of GMM HMMs is almost certainly not "