public class MxTrainer
extends java.lang.Object
implements ai.djl.training.Trainer
MxTrainer is the MXNet implementation of the Trainer.| Modifier and Type | Method and Description |
|---|---|
void |
close() |
void |
endEpoch() |
protected void |
finalize() |
ai.djl.ndarray.NDList |
forward(ai.djl.ndarray.NDList input) |
java.util.List<ai.djl.Device> |
getDevices() |
<T extends ai.djl.training.evaluator.Evaluator> |
getEvaluator(java.lang.Class<T> clazz) |
java.util.List<ai.djl.training.evaluator.Evaluator> |
getEvaluators() |
ai.djl.training.loss.Loss |
getLoss() |
ai.djl.ndarray.NDManager |
getManager() |
ai.djl.metric.Metrics |
getMetrics() |
ai.djl.Model |
getModel() |
void |
initialize(ai.djl.ndarray.types.Shape... shapes) |
ai.djl.training.GradientCollector |
newGradientCollector() |
void |
setMetrics(ai.djl.metric.Metrics metrics) |
void |
step() |
void |
trainBatch(ai.djl.training.dataset.Batch batch) |
void |
validateBatch(ai.djl.training.dataset.Batch batch) |
public void initialize(ai.djl.ndarray.types.Shape... shapes)
initialize in interface ai.djl.training.Trainerpublic ai.djl.training.GradientCollector newGradientCollector()
newGradientCollector in interface ai.djl.training.Trainerpublic void trainBatch(ai.djl.training.dataset.Batch batch)
trainBatch in interface ai.djl.training.Trainerpublic ai.djl.ndarray.NDList forward(ai.djl.ndarray.NDList input)
forward in interface ai.djl.training.Trainerpublic void validateBatch(ai.djl.training.dataset.Batch batch)
validateBatch in interface ai.djl.training.Trainerpublic void step()
step in interface ai.djl.training.Trainerpublic ai.djl.metric.Metrics getMetrics()
getMetrics in interface ai.djl.training.Trainerpublic void setMetrics(ai.djl.metric.Metrics metrics)
setMetrics in interface ai.djl.training.Trainerpublic java.util.List<ai.djl.Device> getDevices()
getDevices in interface ai.djl.training.Trainerpublic void endEpoch()
endEpoch in interface ai.djl.training.Trainerpublic ai.djl.training.loss.Loss getLoss()
getLoss in interface ai.djl.training.Trainerpublic ai.djl.Model getModel()
getModel in interface ai.djl.training.Trainerpublic java.util.List<ai.djl.training.evaluator.Evaluator> getEvaluators()
getEvaluators in interface ai.djl.training.Trainerpublic final <T extends ai.djl.training.evaluator.Evaluator> T getEvaluator(java.lang.Class<T> clazz)
getEvaluator in interface ai.djl.training.Trainerpublic ai.djl.ndarray.NDManager getManager()
getManager in interface ai.djl.training.Trainerprotected void finalize()
throws java.lang.Throwable
finalize in class java.lang.Objectjava.lang.Throwablepublic void close()
close in interface ai.djl.training.Trainerclose in interface java.lang.AutoCloseable