Class CrossEntropyLoss

java.lang.Object
io.github.kirstenali.deepj.loss.CrossEntropyLoss
All Implemented Interfaces:
LossFunction

public final class CrossEntropyLoss extends Object implements LossFunction
Cross-entropy loss with integer class targets.

Expected shapes:

  • predicted (logits): [nTokens x vocab]
  • actual (class indices): [nTokens x 1] where each entry is an integer in [0, vocab)

This class also provides helpers for common language-modeling usage where targets are provided as int[].

  • Constructor Details

    • CrossEntropyLoss

      public CrossEntropyLoss()
  • Method Details

    • loss

      public float loss(Tensor predicted, Tensor actual)
      Specified by:
      loss in interface LossFunction
    • gradient

      public Tensor gradient(Tensor predicted, Tensor actual)
      Specified by:
      gradient in interface LossFunction
    • loss

      public static float loss(Tensor logits, int[] targets)
      Convenience helper: compute loss from logits and int targets.
    • gradient

      public static Tensor gradient(Tensor logits, int[] targets)
      Convenience helper: gradient w.r.t. logits, averaged over rows.
    • toIntTargets

      public static int[] toIntTargets(Tensor actual)
      Converts a [n x 1] Tensor of class indices into an int[].
    • fromIntTargets

      public static Tensor fromIntTargets(int[] targets)
      Builds a [n x 1] Tensor from int[] targets.