Gmm

View as Markdown

Source header: cuvs/cluster/gmm.hpp

Gaussian mixture hyperparameters

cluster::gmm::covariance_type

Covariance parameterization of the mixture components.

1enum class covariance_type {
2 FULL = 0,
3 TIED = 1,
4 DIAG = 2,
5 SPHERICAL = 3
6};

Values

NameValue
FULL0
TIED1
DIAG2
SPHERICAL3

cluster::gmm::init_method

Strategy used to initialize the responsibilities before EM.

1enum class init_method {
2 KMeans = 0,
3 KMeansPlusPlus = 1,
4 Random = 2,
5 RandomFromData = 3
6};

Values

NameValue
KMeans0
KMeansPlusPlus1
Random2
RandomFromData3

cluster::gmm::params

Hyper-parameters for the Gaussian mixture EM solver.

1struct params {
2 int n_components;
3 covariance_type cov_type;
4 double tol;
5 double reg_covar;
6 int max_iter;
7 int n_init;
8 init_method init;
9 uint64_t seed;
10};

Fields

NameTypeDescription
n_componentsintThe number of mixture components (at most 65535). Default: 1.
cov_typecovariance_typeCovariance parameterization of the mixture components. Default: FULL.
toldoubleConvergence threshold on the change of the per-sample average log-likelihood (lower bound). Default: 1e-3.
reg_covardoubleNon-negative regularization added to the diagonal of covariance.
Default: 1e-6.
max_iterintMaximum number of EM iterations for a single run. Default: 100.
n_initintNumber of initializations to perform; the best result is kept.
Default: 1.
initinit_methodStrategy used to initialize the responsibilities before EM.
Default: KMeans.
seeduint64_tSeed to the random number generator. Default: 0.

Gaussian mixture model APIs

cluster::gmm::fit

Fit a Gaussian mixture with the EM algorithm.

1void fit(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const float, int64_t> X,
4raft::device_vector_view<float, int64_t> weights,
5raft::device_matrix_view<float, int64_t> means,
6raft::device_vector_view<float, int64_t> covariances,
7raft::device_vector_view<float, int64_t> precisions_chol,
8raft::device_vector_view<float, int64_t> precisions,
9raft::device_vector_view<int, int64_t> labels,
10raft::host_scalar_view<float> lower_bound,
11raft::host_scalar_view<int> n_iter,
12raft::host_scalar_view<bool> converged,
13bool warm_start = false);

Runs params.n_init random restarts (unless warm_start is true) and keeps the parameters with the largest lower bound. Writes the fitted weights, means, covariances, precisions_chol and precisions, the per-sample hard labels (argmax of the final responsibilities), and the scalar lower_bound / n_iter / converged diagnostics.

When warm_start is true the incoming weights / means / covariances are used as the single initialization and params.n_init is ignored.

Parameters

NameDirectionTypeDescription
handleinraft::resources const&The raft resources handle.
paramsinconst params&Hyper-parameters of the EM solver.
Xinraft::device_matrix_view<const float, int64_t>Training data, row-major. [dim = n_samples x n_features]
weightsinoutraft::device_vector_view<float, int64_t>Mixture weights. [len = n_components]
meansinoutraft::device_matrix_view<float, int64_t>Component means, row-major. [dim = n_components x n_features]
covariancesinoutraft::device_vector_view<float, int64_t>Component covariances, flat. Length depends on cov_type (K=n_components, d=n_features): FULL Kdd, TIED dd, DIAG Kd, SPHERICAL K.
precisions_choloutraft::device_vector_view<float, int64_t>Precision Cholesky factors, same flat layout as covariances. FULL/TIED hold the upper-triangular factor U (precision = U @ Uᵀ); DIAG/SPHERICAL hold reciprocal standard deviations.
precisionsoutraft::device_vector_view<float, int64_t>Precision matrices, same flat layout as covariances.
labelsoutraft::device_vector_view<int, int64_t>Hard component assignment per sample. [len = n_samples]
lower_boundoutraft::host_scalar_view<float>Per-sample average log-likelihood of the best fit.
n_iteroutraft::host_scalar_view<int>Number of EM iterations of the best fit.
convergedoutraft::host_scalar_view<bool>Whether the best fit converged within params.tol.
warm_startinboolUse the incoming weights/means/covariances as the single initialization.
Default: false.

Returns

void

Additional overload: cluster::gmm::fit

1void fit(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const double, int64_t> X,
4raft::device_vector_view<double, int64_t> weights,
5raft::device_matrix_view<double, int64_t> means,
6raft::device_vector_view<double, int64_t> covariances,
7raft::device_vector_view<double, int64_t> precisions_chol,
8raft::device_vector_view<double, int64_t> precisions,
9raft::device_vector_view<int, int64_t> labels,
10raft::host_scalar_view<double> lower_bound,
11raft::host_scalar_view<int> n_iter,
12raft::host_scalar_view<bool> converged,
13bool warm_start = false);

Parameters

NameDirectionTypeDescription
handleraft::resources const&
paramsconst params&
Xraft::device_matrix_view<const double, int64_t>
weightsraft::device_vector_view<double, int64_t>
meansraft::device_matrix_view<double, int64_t>
covariancesraft::device_vector_view<double, int64_t>
precisions_cholraft::device_vector_view<double, int64_t>
precisionsraft::device_vector_view<double, int64_t>
labelsraft::device_vector_view<int, int64_t>
lower_boundraft::host_scalar_view<double>
n_iterraft::host_scalar_view<int>
convergedraft::host_scalar_view<bool>
warm_startboolDefault: false.

Returns

void

cluster::gmm::predict

Hard component labels (argmax responsibility) for new data.

1void predict(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const float, int64_t> X,
4raft::device_vector_view<const float, int64_t> weights,
5raft::device_matrix_view<const float, int64_t> means,
6raft::device_vector_view<const float, int64_t> precisions_chol,
7raft::device_vector_view<int, int64_t> labels);

Parameters

NameDirectionTypeDescription
handleinraft::resources const&The raft resources handle.
paramsinconst params&Fit hyper-parameters; only n_components and cov_type are consulted at inference time.
Xinraft::device_matrix_view<const float, int64_t>Data to assign, row-major. [dim = n_samples x n_features]
weightsinraft::device_vector_view<const float, int64_t>Fitted mixture weights. [len = n_components]
meansinraft::device_matrix_view<const float, int64_t>Fitted component means. [dim = n_components x n_features]
precisions_cholinraft::device_vector_view<const float, int64_t>Fitted precision Cholesky factors, flat. Length by cov_type (K=n_components, d=n_features): FULL Kdd, TIED dd, DIAG Kd, SPHERICAL K.
labelsoutraft::device_vector_view<int, int64_t>Hard component assignment per sample. [len = n_samples]

Returns

void

Additional overload: cluster::gmm::predict

1void predict(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const double, int64_t> X,
4raft::device_vector_view<const double, int64_t> weights,
5raft::device_matrix_view<const double, int64_t> means,
6raft::device_vector_view<const double, int64_t> precisions_chol,
7raft::device_vector_view<int, int64_t> labels);

Parameters

NameDirectionTypeDescription
handleraft::resources const&
paramsconst params&
Xraft::device_matrix_view<const double, int64_t>
weightsraft::device_vector_view<const double, int64_t>
meansraft::device_matrix_view<const double, int64_t>
precisions_cholraft::device_vector_view<const double, int64_t>
labelsraft::device_vector_view<int, int64_t>

Returns

void

cluster::gmm::predict_proba

Posterior responsibilities for new data.

1void predict_proba(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const float, int64_t> X,
4raft::device_vector_view<const float, int64_t> weights,
5raft::device_matrix_view<const float, int64_t> means,
6raft::device_vector_view<const float, int64_t> precisions_chol,
7raft::device_matrix_view<float, int64_t> resp);

Parameters

NameDirectionTypeDescription
handleinraft::resources const&The raft resources handle.
paramsinconst params&Fit hyper-parameters; only n_components and cov_type are consulted at inference time.
Xinraft::device_matrix_view<const float, int64_t>Data to evaluate, row-major. [dim = n_samples x n_features]
weightsinraft::device_vector_view<const float, int64_t>Fitted mixture weights. [len = n_components]
meansinraft::device_matrix_view<const float, int64_t>Fitted component means. [dim = n_components x n_features]
precisions_cholinraft::device_vector_view<const float, int64_t>Fitted precision Cholesky factors, flat. Length by cov_type (K=n_components, d=n_features): FULL Kdd, TIED dd, DIAG Kd, SPHERICAL K.
respoutraft::device_matrix_view<float, int64_t>Posterior probability of each component for each sample, row-major. [dim = n_samples x n_components]

Returns

void

Additional overload: cluster::gmm::predict_proba

1void predict_proba(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const double, int64_t> X,
4raft::device_vector_view<const double, int64_t> weights,
5raft::device_matrix_view<const double, int64_t> means,
6raft::device_vector_view<const double, int64_t> precisions_chol,
7raft::device_matrix_view<double, int64_t> resp);

Parameters

NameDirectionTypeDescription
handleraft::resources const&
paramsconst params&
Xraft::device_matrix_view<const double, int64_t>
weightsraft::device_vector_view<const double, int64_t>
meansraft::device_matrix_view<const double, int64_t>
precisions_cholraft::device_vector_view<const double, int64_t>
respraft::device_matrix_view<double, int64_t>

Returns

void

cluster::gmm::score_samples

Per-sample log-likelihood log p(x_i) for new data.

1void score_samples(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const float, int64_t> X,
4raft::device_vector_view<const float, int64_t> weights,
5raft::device_matrix_view<const float, int64_t> means,
6raft::device_vector_view<const float, int64_t> precisions_chol,
7raft::device_vector_view<float, int64_t> log_prob_norm);

Parameters

NameDirectionTypeDescription
handleinraft::resources const&The raft resources handle.
paramsinconst params&Fit hyper-parameters; only n_components and cov_type are consulted at inference time.
Xinraft::device_matrix_view<const float, int64_t>Data to evaluate, row-major. [dim = n_samples x n_features]
weightsinraft::device_vector_view<const float, int64_t>Fitted mixture weights. [len = n_components]
meansinraft::device_matrix_view<const float, int64_t>Fitted component means. [dim = n_components x n_features]
precisions_cholinraft::device_vector_view<const float, int64_t>Fitted precision Cholesky factors, flat. Length by cov_type (K=n_components, d=n_features): FULL Kdd, TIED dd, DIAG Kd, SPHERICAL K.
log_prob_normoutraft::device_vector_view<float, int64_t>Log-likelihood of each sample under the model. [len = n_samples]

Returns

void

Additional overload: cluster::gmm::score_samples

1void score_samples(raft::resources const& handle,
2const params& params,
3raft::device_matrix_view<const double, int64_t> X,
4raft::device_vector_view<const double, int64_t> weights,
5raft::device_matrix_view<const double, int64_t> means,
6raft::device_vector_view<const double, int64_t> precisions_chol,
7raft::device_vector_view<double, int64_t> log_prob_norm);

Parameters

NameDirectionTypeDescription
handleraft::resources const&
paramsconst params&
Xraft::device_matrix_view<const double, int64_t>
weightsraft::device_vector_view<const double, int64_t>
meansraft::device_matrix_view<const double, int64_t>
precisions_cholraft::device_vector_view<const double, int64_t>
log_prob_normraft::device_vector_view<double, int64_t>

Returns

void