Class MxGradientCollector

java.lang.Object
ai.djl.mxnet.engine.MxGradientCollector
All Implemented Interfaces:
ai.djl.training.GradientCollector, AutoCloseable

public final class MxGradientCollector extends Object implements ai.djl.training.GradientCollector
MxGradientCollector is the MXNet implementation of GradientCollector.
  • Method Summary

    Modifier and Type
    Method
    Description
    void
    backward(ai.djl.ndarray.NDArray array)
    void
    static Symbol
    getSymbol(ai.djl.ndarray.NDManager manager, ai.djl.ndarray.NDArray array)
    Returns the Symbol of a network formed by the recorded operations on the given NDArray.
    static boolean
    Gets whether Autograd is recording computations.
    static boolean
    Gets whether Autograd is in training/predicting mode.
    static boolean
    setRecording(boolean isRecording)
    Sets the status to recording/not recording.
    static boolean
    setTraining(boolean isTraining)
    Sets the status to training/predicting.
    void

    Methods inherited from class java.lang.Object

    clone, equals, finalize, getClass, hashCode, notify, notifyAll, toString, wait, wait, wait
  • Method Details

    • isRecording

      public static boolean isRecording()
      Gets whether Autograd is recording computations.
      Returns:
      the current state of recording
    • isTraining

      public static boolean isTraining()
      Gets whether Autograd is in training/predicting mode.
      Returns:
      the current state of training/predicting
    • setRecording

      public static boolean setRecording(boolean isRecording)
      Sets the status to recording/not recording. When recording, graph will be constructed for gradient computation.
      Parameters:
      isRecording - the recording state to be set
      Returns:
      the previous recording state before this set
    • setTraining

      public static boolean setTraining(boolean isTraining)
      Sets the status to training/predicting. This affects ctx.is_train in the device running the operator. For example, Dropout will drop inputs randomly when isTraining=True, while simply passing through if isTraining=False.
      Parameters:
      isTraining - true if for training
      Returns:
      the previous status before this set
    • getSymbol

      public static Symbol getSymbol(ai.djl.ndarray.NDManager manager, ai.djl.ndarray.NDArray array)
      Returns the Symbol of a network formed by the recorded operations on the given NDArray.
      Parameters:
      manager - the NDManager to create the Symbol
      array - the NDArray
      Returns:
      the Symbol
    • close

      public void close()
      Specified by:
      close in interface AutoCloseable
      Specified by:
      close in interface ai.djl.training.GradientCollector
    • backward

      public void backward(ai.djl.ndarray.NDArray array)
      Specified by:
      backward in interface ai.djl.training.GradientCollector
    • zeroGradients

      public void zeroGradients()
      Specified by:
      zeroGradients in interface ai.djl.training.GradientCollector