Package ai.djl.mxnet.engine
Class MxGradientCollector
java.lang.Object
ai.djl.mxnet.engine.MxGradientCollector
- All Implemented Interfaces:
ai.djl.training.GradientCollector,AutoCloseable
MxGradientCollector is the MXNet implementation of GradientCollector.-
Method Summary
Modifier and TypeMethodDescriptionvoidbackward(ai.djl.ndarray.NDArray array) voidclose()static SymbolgetSymbol(ai.djl.ndarray.NDManager manager, ai.djl.ndarray.NDArray array) Returns theSymbolof a network formed by the recorded operations on the givenNDArray.static booleanGets whether Autograd is recording computations.static booleanGets whether Autograd is in training/predicting mode.static booleansetRecording(boolean isRecording) Sets the status to recording/not recording.static booleansetTraining(boolean isTraining) Sets the status to training/predicting.void
-
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-trueif for training- Returns:
- the previous status before this set
-
getSymbol
Returns theSymbolof a network formed by the recorded operations on the givenNDArray. -
close
public void close()- Specified by:
closein interfaceAutoCloseable- Specified by:
closein interfaceai.djl.training.GradientCollector
-
backward
public void backward(ai.djl.ndarray.NDArray array) - Specified by:
backwardin interfaceai.djl.training.GradientCollector
-
zeroGradients
public void zeroGradients()- Specified by:
zeroGradientsin interfaceai.djl.training.GradientCollector
-