Class BinaryCrossEntropyLoss

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

public class BinaryCrossEntropyLoss extends Object implements LossFunction, Serializable
Cross Entropy Loss is a loss function used for binary classification tasks (two classes, single output which represents probability ). It should be used in combination with sigmoid output activation function. The formula: E = (1/n) * -SUM( t * ln(y) + (1-t) * ln(1-y) ) where t is target, and y actual output Bishop, C. pg. 231, eq. 6.120
See Also:
  • Constructor Details

    • BinaryCrossEntropyLoss

      public BinaryCrossEntropyLoss(NeuralNetwork neuralNet)
  • Method Details

    • addPatternError

      public TensorBase addPatternError(TensorBase predictedOutput, TensorBase targetOutput)
      Calculates error for given actual and target patterns and adds that error to total error. Returns output error vector for specified actual and target outputs.
      Specified by:
      addPatternError in interface LossFunction
      Parameters:
      predictedOutput - predicted output of a neural network
      targetOutput - target output of a neural network
      Returns:
      error vector for specified predicted 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