Class CrossEntropyLoss

java.lang.Object
deepnetts.net.loss.CrossEntropyLoss
All Implemented Interfaces:
LossFunction, Serializable

public class CrossEntropyLoss extends Object implements LossFunction, Serializable
Average Cross Entropy Loss function commonly used for multi class classification problems. E = -1/n * SUM(SUM(t*ln(y))) Since its 1-of-n classification scheme, all outputs except target are zeros so it comes down to E = -1/n * SUM(ln(y_targetIdx))
See Also:
  • Constructor Details

    • CrossEntropyLoss

      public CrossEntropyLoss(NeuralNetwork neuralNet)
  • Method Details

    • addPatternError

      public TensorBase addPatternError(TensorBase predictedOut, TensorBase targetOut)
      Calculates and returns outpurt error vector for specified predicted and target outputs.
      Specified by:
      addPatternError in interface LossFunction
      Parameters:
      predictedOut - predicted output from the neural network
      targetOut - target/desired output of the neural network
      Returns:
      error vector for specified actual and target outputs
    • getPatternLoss

      public float getPatternLoss()
      Specified by:
      getPatternLoss in interface LossFunction
    • addRegularizationSum

      public void addRegularizationSum(float regSum)
      Description copied from interface: LossFunction
      Adds specified regularization sum to total loss.
      Specified by:
      addRegularizationSum in interface LossFunction
      Parameters:
      regSum - regularization sum
    • getTotal

      public float getTotal()
      Description copied from interface: LossFunction
      Returns the total error calculated by this loss function.
      Specified by:
      getTotal in interface LossFunction
      Returns:
      total error calculated by this loss function
    • reset

      public void reset()
      Description copied from interface: LossFunction
      Resets the total error and pattern counter.
      Specified by:
      reset in interface LossFunction