diff --git a/src/mlpack/methods/lmnn/lmnn.hpp b/src/mlpack/methods/lmnn/lmnn.hpp index 0517d5ec51..a62f8b86e4 100644 --- a/src/mlpack/methods/lmnn/lmnn.hpp +++ b/src/mlpack/methods/lmnn/lmnn.hpp @@ -79,7 +79,8 @@ class LMNN * * @param outputMatrix Covariance matrix of Mahalanobis distance. */ - void LearnDistance(arma::mat& outputMatrix); + template + void LearnDistance(arma::mat& outputMatrix, CallbackTypes&&... callbacks); //! Get the dataset reference. diff --git a/src/mlpack/methods/lmnn/lmnn_impl.hpp b/src/mlpack/methods/lmnn/lmnn_impl.hpp index 3f6b539bee..4669e74281 100644 --- a/src/mlpack/methods/lmnn/lmnn_impl.hpp +++ b/src/mlpack/methods/lmnn/lmnn_impl.hpp @@ -36,7 +36,9 @@ LMNN::LMNN(const arma::mat& dataset, { /* nothing to do */ } template -void LMNN::LearnDistance(arma::mat& outputMatrix) +template +void LMNN::LearnDistance(arma::mat& outputMatrix, + CallbackTypes&&... callbacks) { // LMNN objective function. LMNNFunction objFunction(dataset, labels, k, @@ -56,7 +58,7 @@ void LMNN::LearnDistance(arma::mat& outputMatrix) Timer::Start("lmnn_optimization"); - optimizer.Optimize(objFunction, outputMatrix); + optimizer.Optimize(objFunction, outputMatrix, callbacks...); Timer::Stop("lmnn_optimization"); }