Rename weightVectors to weights and simplify API.
This commit is contained in:
@@ -30,26 +30,30 @@ class SimpleWeightUpdate
|
||||
* the weights of the incorrectly classified class while increasing the weight
|
||||
* of the correct class it should have been classified to.
|
||||
*
|
||||
* @param trainData The training dataset.
|
||||
* @param weightVectors Matrix of weight vectors.
|
||||
* @param colIndex Index of the column which has been incorrectly predicted.
|
||||
* @param labelIndex Index of the vector in trainData.
|
||||
* @param vectorIndex Index of the class which should have been predicted.
|
||||
* @param D Cost of mispredicting the labelIndex instance.
|
||||
* @tparam Type of vector (should be an Armadillo vector like arma::vec or
|
||||
* arma::sp_vec or something similar).
|
||||
* @param trainingPoint Point that was misclassified.
|
||||
* @param weights Matrix of weights.
|
||||
* @param biases Vector of biases.
|
||||
* @param incorrectClass Index of class that the point was incorrectly
|
||||
* classified as.
|
||||
* @param correctClass Index of the true class of the point.
|
||||
* @param instanceWeight Weight to be given to this particular point during
|
||||
* training (this is useful for boosting).
|
||||
*/
|
||||
void UpdateWeights(const arma::mat& trainData,
|
||||
arma::mat& weightVectors,
|
||||
template<typename VecType>
|
||||
void UpdateWeights(const VecType& trainingPoint,
|
||||
arma::mat& weights,
|
||||
arma::vec& biases,
|
||||
const size_t labelIndex,
|
||||
const size_t vectorIndex,
|
||||
const size_t colIndex,
|
||||
const arma::rowvec& D)
|
||||
const size_t incorrectClass,
|
||||
const size_t correctClass,
|
||||
const double instanceWeight = 1.0)
|
||||
{
|
||||
weightVectors.col(colIndex) -= D(labelIndex) * trainData.col(labelIndex);
|
||||
biases(colIndex) -= D(labelIndex);
|
||||
weights.col(incorrectClass) -= instanceWeight * trainingPoint;
|
||||
biases(incorrectClass) -= instanceWeight;
|
||||
|
||||
weightVectors.col(vectorIndex) += D(labelIndex) * trainData.col(labelIndex);
|
||||
biases(vectorIndex) += D(labelIndex);
|
||||
weights.col(correctClass) += instanceWeight * trainingPoint;
|
||||
biases(correctClass) += instanceWeight;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ class Perceptron
|
||||
{
|
||||
public:
|
||||
/**
|
||||
* Constructor - constructs the perceptron by building the weightVectors
|
||||
* Constructor - constructs the perceptron by building the weights
|
||||
* matrix, which is later used in Classification. It adds a bias input vector
|
||||
* of 1 to the input data to take care of the bias weights.
|
||||
*
|
||||
@@ -46,7 +46,7 @@ class Perceptron
|
||||
const int iterations);
|
||||
|
||||
/**
|
||||
* Classification function. After training, use the weightVectors matrix to
|
||||
* Classification function. After training, use the weights matrix to
|
||||
* classify test, and put the predicted classes in predictedLabels.
|
||||
*
|
||||
* @param test Testing data or data to classify.
|
||||
@@ -81,12 +81,12 @@ private:
|
||||
size_t iter;
|
||||
|
||||
/**
|
||||
* Stores the weight vectors for each of the input class labels. Each column
|
||||
* Stores the weights for each of the input class labels. Each column
|
||||
* corresponds to the weights for one class label, and each row corresponds to
|
||||
* the weights for one dimension of the input data. The biases are held in a
|
||||
* separate vector.
|
||||
*/
|
||||
arma::mat weightVectors;
|
||||
arma::mat weights;
|
||||
|
||||
//! The biases for each class.
|
||||
arma::vec biases;
|
||||
|
||||
@@ -13,10 +13,9 @@ namespace mlpack {
|
||||
namespace perceptron {
|
||||
|
||||
/**
|
||||
* Constructor - constructs the perceptron. Or rather, builds the weightVectors
|
||||
* matrix, which is later used in Classification.
|
||||
* It adds a bias input vector of 1 to the input data to take care of the bias
|
||||
* weights.
|
||||
* Constructor - constructs the perceptron. Or rather, builds the weights
|
||||
* matrix, which is later used in classification. It adds a bias input vector
|
||||
* of 1 to the input data to take care of the bias weights.
|
||||
*
|
||||
* @param data Input, training data.
|
||||
* @param labels Labels of dataset.
|
||||
@@ -34,7 +33,7 @@ Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Perceptron(
|
||||
const int iterations)
|
||||
{
|
||||
WeightInitializationPolicy WIP;
|
||||
WIP.Initialize(weightVectors, biases, data.n_rows, arma::max(labels) + 1);
|
||||
WIP.Initialize(weights, biases, data.n_rows, arma::max(labels) + 1);
|
||||
|
||||
// Start training.
|
||||
iter = iterations;
|
||||
@@ -46,8 +45,8 @@ Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Perceptron(
|
||||
|
||||
|
||||
/**
|
||||
* Classification function. After training, use the weightVectors matrix to
|
||||
* classify test, and put the predicted classes in predictedLabels.
|
||||
* Classification function. After training, use the weights matrix to classify
|
||||
* test, and put the predicted classes in predictedLabels.
|
||||
*
|
||||
* @param test testing data or data to classify.
|
||||
* @param predictedLabels vector to store the predicted classes after
|
||||
@@ -68,7 +67,7 @@ void Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Classify(
|
||||
// Could probably be faster if done in batch.
|
||||
for (size_t i = 0; i < test.n_cols; i++)
|
||||
{
|
||||
tempLabelMat = weightVectors.t() * test.col(i) + biases;
|
||||
tempLabelMat = weights.t() * test.col(i) + biases;
|
||||
tempLabelMat.max(maxIndex);
|
||||
predictedLabels(0, i) = maxIndex;
|
||||
}
|
||||
@@ -99,7 +98,7 @@ Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Perceptron(
|
||||
|
||||
// Insert a row of ones at the top of the training data set.
|
||||
WeightInitializationPolicy WIP;
|
||||
WIP.Initialize(weightVectors, biases, data.n_rows, arma::max(labels) + 1);
|
||||
WIP.Initialize(weights, biases, data.n_rows, arma::max(labels) + 1);
|
||||
|
||||
Train(data, labels, D);
|
||||
}
|
||||
@@ -151,7 +150,7 @@ void Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Train(
|
||||
{
|
||||
// Multiply for each variable and check whether the current weight vector
|
||||
// correctly classifies this.
|
||||
tempLabelMat = weightVectors.t() * data.col(j) + biases;
|
||||
tempLabelMat = weights.t() * data.col(j) + biases;
|
||||
|
||||
tempLabelMat.max(maxIndexRow, maxIndexCol);
|
||||
|
||||
@@ -164,8 +163,8 @@ void Perceptron<LearnPolicy, WeightInitializationPolicy, MatType>::Train(
|
||||
// Send maxIndexRow for knowing which weight to update, send j to know
|
||||
// the value of the vector to update it with. Send tempLabel to know
|
||||
// the correct class.
|
||||
LP.UpdateWeights(data, weightVectors, biases, j, tempLabel, maxIndexRow,
|
||||
D);
|
||||
LP.UpdateWeights(data.col(j), weights, biases, maxIndexRow, tempLabel,
|
||||
D(j));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user