Use Ryan's strategy to handle variable batch size
This commit is contained in:
@@ -226,9 +226,12 @@ class LMNNFunction
|
||||
arma::mat maxImpNorm;
|
||||
//! Holds previous transformation matrix. Used for L-BFGS like optimizer.
|
||||
arma::mat transformationOld;
|
||||
//! Holds previous transformation matrix for each point. Used for
|
||||
//! optimizers which operates over batches.
|
||||
arma::cube transformationOldPoint;
|
||||
//! Holds previous transformation matrices.
|
||||
std::vector<arma::mat> oldTransformationMatrices;
|
||||
//! Holds number of points which are using each transformation matrix.
|
||||
std::vector<size_t> oldTransformationCounts;
|
||||
//! Holds points to transformation matrix mapping.
|
||||
arma::vec lastTransformationIndices;
|
||||
/**
|
||||
* Precalculate the gradient part due to target neighbors and stores
|
||||
* the result as a matrix. Used for L-BFGS like optimizers which does not
|
||||
|
||||
@@ -41,12 +41,16 @@ LMNNFunction<MetricType>::LMNNFunction(const arma::mat& dataset,
|
||||
// Initialize transformed dataset to base dataset.
|
||||
transformedDataset = dataset;
|
||||
|
||||
// Initialize cache.
|
||||
evalOld.set_size(k, k, dataset.n_cols);
|
||||
evalOld.fill(arma::datum::nan);
|
||||
|
||||
maxImpNorm.set_size(k, dataset.n_cols);
|
||||
maxImpNorm.fill(0);
|
||||
|
||||
lastTransformationIndices.set_size(dataset.n_cols);
|
||||
lastTransformationIndices.fill(arma::datum::nan);
|
||||
|
||||
// Initialize target neighbors & impostors.
|
||||
targetNeighbors = arma::Mat<size_t>(k, dataset.n_cols, arma::fill::zeros);
|
||||
impostors = arma::Mat<size_t>(k, dataset.n_cols, arma::fill::zeros);
|
||||
@@ -66,7 +70,7 @@ void LMNNFunction<MetricType>::Shuffle()
|
||||
arma::mat newDataset = dataset;
|
||||
arma::Mat<size_t> newLabels = labels;
|
||||
arma::cube newEvalOld = evalOld;
|
||||
arma::cube newTransformationOldPoint = transformationOldPoint;
|
||||
arma::vec newlastTransformationIndices = lastTransformationIndices;
|
||||
arma::mat newMaxImpNorm = maxImpNorm;
|
||||
|
||||
// Generate ordering.
|
||||
@@ -79,11 +83,11 @@ void LMNNFunction<MetricType>::Shuffle()
|
||||
dataset = newDataset.cols(ordering);
|
||||
labels = newLabels.cols(ordering);
|
||||
maxImpNorm = newMaxImpNorm.cols(ordering);
|
||||
lastTransformationIndices = newlastTransformationIndices.elem(ordering);
|
||||
|
||||
for (size_t i = 0; i < ordering.n_elem; i++)
|
||||
{
|
||||
evalOld.slice(i) = newEvalOld.slice(ordering(i));
|
||||
transformationOldPoint.slice(i) = newTransformationOldPoint.slice(ordering(i));
|
||||
}
|
||||
|
||||
// Re-calculate target neighbors as indices changed.
|
||||
@@ -207,14 +211,6 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
{
|
||||
double cost = 0;
|
||||
|
||||
// Ensure cache transformation cube has correct size.
|
||||
if (transformationOldPoint.n_elem == 0)
|
||||
{
|
||||
transformationOldPoint.set_size(transformation.n_rows,
|
||||
transformation.n_cols, dataset.n_cols);
|
||||
transformationOldPoint.fill(arma::datum::nan);
|
||||
}
|
||||
|
||||
// Apply metric over dataset.
|
||||
transformedDataset = transformation * dataset;
|
||||
|
||||
@@ -237,10 +233,10 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
|
||||
// Calculate norm of change in transformation.
|
||||
double transformationDiff = 0;
|
||||
if (!transformationOldPoint.slice(i).has_nan())
|
||||
if (arma::is_finite(lastTransformationIndices(i)))
|
||||
{
|
||||
transformationDiff = arma::norm(transformation -
|
||||
transformationOldPoint.slice(i));
|
||||
oldTransformationMatrices[lastTransformationIndices(i)]);
|
||||
}
|
||||
|
||||
for (int j = k - 1; j >= 0; j--)
|
||||
@@ -257,7 +253,7 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
bool exactEval = true;
|
||||
|
||||
// Bounds for eval.
|
||||
if (transformationOld.n_elem != 0 && !std::isnan(evalOld(l, j, i)))
|
||||
if (arma::is_finite(lastTransformationIndices(i)) && !std::isnan(evalOld(l, j, i)))
|
||||
{
|
||||
// Update cache max impostor norm.
|
||||
maxImpNorm(l, i) = std::max(maxImpNorm(l, i), norm(impostors(l, i)));
|
||||
@@ -268,15 +264,7 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
|
||||
// Check if there is need to calculate exact eval value.
|
||||
if (eval <= -1)
|
||||
{
|
||||
exactEval = false;
|
||||
}
|
||||
else
|
||||
{
|
||||
// Reset cacche max impostor norm.
|
||||
maxImpNorm(l, i) = 0;
|
||||
evalOld(l, j, i) = arma::datum::nan;
|
||||
}
|
||||
}
|
||||
|
||||
// Calculate exact eval value.
|
||||
@@ -300,9 +288,6 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
// Update cache eval value.
|
||||
evalOld(l, j, i) = eval;
|
||||
|
||||
// Update cache transformation matrix.
|
||||
transformationOldPoint.slice(i) = transformation;
|
||||
|
||||
// Check bounding condition.
|
||||
if (eval <= -1)
|
||||
{
|
||||
@@ -312,10 +297,52 @@ double LMNNFunction<MetricType>::Evaluate(const arma::mat& transformation,
|
||||
}
|
||||
|
||||
cost += regularization * (1 + eval);
|
||||
|
||||
// Reset cache.
|
||||
if (eval > -1 && arma::is_finite(lastTransformationIndices(i)))
|
||||
{
|
||||
// update bound.
|
||||
evalOld(l, j, i) = arma::datum::nan;
|
||||
maxImpNorm(l, i)=0;
|
||||
lastTransformationIndices(i) = arma::datum::nan;
|
||||
--oldTransformationCounts[lastTransformationIndices(i)];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Update cache transformation matrices.
|
||||
// Are there any empty transformation matrices?
|
||||
size_t index = oldTransformationMatrices.size();
|
||||
for (size_t i = 0; i < oldTransformationCounts.size(); ++i)
|
||||
{
|
||||
if (oldTransformationCounts[i] == 0)
|
||||
{
|
||||
index = i; // Reuse this index.
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Did we find an unused matrix? If not, we have to allocate new space.
|
||||
if (index == oldTransformationMatrices.size())
|
||||
{
|
||||
oldTransformationMatrices.push_back(transformation);
|
||||
oldTransformationCounts.push_back(0);
|
||||
}
|
||||
else
|
||||
{
|
||||
oldTransformationMatrices[index] = transformation;
|
||||
}
|
||||
|
||||
// Update all the transformation indices.
|
||||
for (size_t i = begin; i < begin + batchSize; ++i)
|
||||
{
|
||||
--oldTransformationCounts[lastTransformationIndices(i)];
|
||||
lastTransformationIndices(i) = index;
|
||||
}
|
||||
|
||||
oldTransformationCounts[index] += batchSize;
|
||||
|
||||
return cost;
|
||||
}
|
||||
|
||||
@@ -341,25 +368,10 @@ void LMNNFunction<MetricType>::Gradient(const arma::mat& transformation,
|
||||
for (size_t l = 0, bp = k; l < bp ; l++)
|
||||
{
|
||||
// Calculate gradient due to triplets.
|
||||
double eval = 0;
|
||||
|
||||
// Flag to trigger exact eval calculation.
|
||||
bool exactEval = true;
|
||||
|
||||
// Use eval calualated during Evaluate.
|
||||
if (!std::isnan(evalOld(l, j, i)))
|
||||
{
|
||||
eval = evalOld(l, j, i);
|
||||
exactEval = false;
|
||||
}
|
||||
|
||||
if (exactEval)
|
||||
{
|
||||
eval = metric.Evaluate(transformedDataset.col(i),
|
||||
double eval = metric.Evaluate(transformedDataset.col(i),
|
||||
transformedDataset.col(targetNeighbors(j, i))) -
|
||||
metric.Evaluate(transformedDataset.col(i),
|
||||
transformedDataset.col(impostors(l, i)));
|
||||
}
|
||||
|
||||
// Check bounding condition.
|
||||
if (eval < -1)
|
||||
@@ -411,25 +423,10 @@ void LMNNFunction<MetricType>::Gradient(const arma::mat& transformation,
|
||||
for (size_t l = 0, bp = k; l < bp ; l++)
|
||||
{
|
||||
// Calculate gradient due to triplets.
|
||||
double eval = 0;
|
||||
|
||||
// Flag to trigger exact eval calculation.
|
||||
bool exactEval = true;
|
||||
|
||||
// Use eval calualated during Evaluate.
|
||||
if (!std::isnan(evalOld(l, j, i)))
|
||||
{
|
||||
eval = evalOld(l, j, i);
|
||||
exactEval = false;
|
||||
}
|
||||
|
||||
if (exactEval)
|
||||
{
|
||||
eval = metric.Evaluate(transformedDataset.col(i),
|
||||
double eval = metric.Evaluate(transformedDataset.col(i),
|
||||
transformedDataset.col(targetNeighbors(j, i))) -
|
||||
metric.Evaluate(transformedDataset.col(i),
|
||||
transformedDataset.col(impostors(l, i)));
|
||||
}
|
||||
|
||||
// Check bounding condition.
|
||||
if (eval < -1)
|
||||
@@ -592,14 +589,6 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
{
|
||||
double cost = 0;
|
||||
|
||||
// Ensure cache transformation cube has correct size.
|
||||
if (transformationOldPoint.n_elem == 0)
|
||||
{
|
||||
transformationOldPoint.set_size(transformation.n_rows,
|
||||
transformation.n_cols, dataset.n_cols);
|
||||
transformationOldPoint.fill(arma::datum::nan);
|
||||
}
|
||||
|
||||
// Apply metric over dataset.
|
||||
transformedDataset = transformation * dataset;
|
||||
|
||||
@@ -631,10 +620,10 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
|
||||
// Calculate norm of change in transformation.
|
||||
double transformationDiff = 0;
|
||||
if (!transformationOldPoint.slice(i).has_nan())
|
||||
if (arma::is_finite(lastTransformationIndices(i)))
|
||||
{
|
||||
transformationDiff = arma::norm(transformation -
|
||||
transformationOldPoint.slice(i));
|
||||
oldTransformationMatrices[lastTransformationIndices(i)]);
|
||||
}
|
||||
|
||||
for (int j = k - 1; j >= 0; j--)
|
||||
@@ -650,7 +639,7 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
bool exactEval = true;
|
||||
|
||||
// Bounds for eval.
|
||||
if (transformationOld.n_elem != 0 && !std::isnan(evalOld(l, j, i)))
|
||||
if (arma::is_finite(lastTransformationIndices(i)) && !std::isnan(evalOld(l, j, i)))
|
||||
{
|
||||
// Update cache max impostor norm.
|
||||
maxImpNorm(l, i) = std::max(maxImpNorm(l, i), norm(impostors(l, i)));
|
||||
@@ -661,15 +650,7 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
|
||||
// Check if there is need to calculate exact eval value.
|
||||
if (eval <= -1)
|
||||
{
|
||||
exactEval = false;
|
||||
}
|
||||
else
|
||||
{
|
||||
// Reset cacche max impostor norm.
|
||||
maxImpNorm(l, i) = 0;
|
||||
evalOld(l, j, i) = arma::datum::nan;
|
||||
}
|
||||
}
|
||||
|
||||
// Calculate exact eval value.
|
||||
@@ -693,9 +674,6 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
// Update cache eval value.
|
||||
evalOld(l, j, i) = eval;
|
||||
|
||||
// Update cache transformation matrix.
|
||||
transformationOldPoint.slice(i) = transformation;
|
||||
|
||||
// Check bounding condition.
|
||||
if (eval <= -1)
|
||||
{
|
||||
@@ -706,6 +684,16 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
|
||||
cost += regularization * (1 + eval);
|
||||
|
||||
// Reset cache.
|
||||
if (eval > -1 && arma::is_finite(lastTransformationIndices(i)))
|
||||
{
|
||||
// update bound.
|
||||
evalOld(l, j, i) = arma::datum::nan;
|
||||
maxImpNorm(l, i)=0;
|
||||
lastTransformationIndices(i) = arma::datum::nan;
|
||||
--oldTransformationCounts[lastTransformationIndices(i)];
|
||||
}
|
||||
|
||||
// Caculate gradient due to impostors.
|
||||
arma::vec diff = dataset.col(i) - dataset.col(targetNeighbors(j, i));
|
||||
cil += diff * arma::trans(diff);
|
||||
@@ -719,6 +707,38 @@ double LMNNFunction<MetricType>::EvaluateWithGradient(
|
||||
gradient = 2 * transformation * ((1 - regularization) * cij +
|
||||
regularization * cil);
|
||||
|
||||
// Update cache transformation matrices.
|
||||
// Are there any empty transformation matrices?
|
||||
size_t index = oldTransformationMatrices.size();
|
||||
for (size_t i = 0; i < oldTransformationCounts.size(); ++i)
|
||||
{
|
||||
if (oldTransformationCounts[i] == 0)
|
||||
{
|
||||
index = i; // Reuse this index.
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Did we find an unused matrix? If not, we have to allocate new space.
|
||||
if (index == oldTransformationMatrices.size())
|
||||
{
|
||||
oldTransformationMatrices.push_back(transformation);
|
||||
oldTransformationCounts.push_back(0);
|
||||
}
|
||||
else
|
||||
{
|
||||
oldTransformationMatrices[index] = transformation;
|
||||
}
|
||||
|
||||
// Update all the transformation indices.
|
||||
for (size_t i = begin; i < begin + batchSize; ++i)
|
||||
{
|
||||
--oldTransformationCounts[lastTransformationIndices(i)];
|
||||
lastTransformationIndices(i) = index;
|
||||
}
|
||||
|
||||
oldTransformationCounts[index] += batchSize;
|
||||
|
||||
|
||||
return cost;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user