public class Trainer
extends java.lang.Object
implements java.lang.AutoCloseable
Trainer interface provides a session for model training.
Trainer provides an easy, and manageable interface for training. Trainer is
not thread-safe.
See the tutorials on:
| Constructor and Description |
|---|
Trainer(Model model,
TrainingConfig trainingConfig)
|
| Modifier and Type | Method and Description |
|---|---|
void |
close() |
void |
endEpoch()
Runs the end epoch actions.
|
protected void |
finalize() |
NDList |
forward(NDList input)
Applies the forward function of the model once on the given input
NDList. |
java.util.List<Device> |
getDevices()
Returns the devices used for training.
|
java.util.List<Evaluator> |
getEvaluators()
Gets all
Evaluators. |
Loss |
getLoss()
Gets the training
Loss function of the trainer. |
NDManager |
getManager()
Gets the
NDManager from the model. |
Metrics |
getMetrics()
Returns the Metrics param used for benchmarking.
|
Model |
getModel()
Returns the model used to create this trainer.
|
TrainingResult |
getTrainingResult()
Returns the
TrainingResult. |
void |
initialize(Shape... shapes)
Initializes the
Model that the Trainer is going to train. |
java.lang.Iterable<Batch> |
iterateDataset(Dataset dataset)
Fetches an iterator that can iterate through the given
Dataset. |
GradientCollector |
newGradientCollector()
Returns a new instance of
GradientCollector. |
NDList |
predict(NDList input)
Applies the predict function of the model once on the given input
NDList. |
void |
setMetrics(Metrics metrics)
Attaches a Metrics param to use for benchmarking.
|
void |
step()
Updates all of the parameters of the model once.
|
void |
trainBatch(Batch batch)
Trains the model with one iteration of the given
Batch of data. |
void |
validateBatch(Batch batch)
Validates the given batch of data.
|
public Trainer(Model model, TrainingConfig trainingConfig)
model - the model the trainer will train ontrainingConfig - the configuration used by the trainerpublic void initialize(Shape... shapes)
Model that the Trainer is going to train.shapes - an array of Shape of the inputspublic java.lang.Iterable<Batch> iterateDataset(Dataset dataset)
Dataset.dataset - the dataset to iterate throughIterable of Batch that contains batches of data from the datasetpublic GradientCollector newGradientCollector()
GradientCollector.GradientCollectorpublic void trainBatch(Batch batch)
Batch of data.batch - a Batch that contains data, and its respective labelsjava.lang.IllegalArgumentException - if the batch engine does not match the trainer enginepublic NDList forward(NDList input)
NDList.input - the input NDListpublic NDList predict(NDList input)
NDList.input - the input NDListpublic void validateBatch(Batch batch)
During validation, the evaluators and losses are computed, but gradients aren't computed, and parameters aren't updated.
batch - a Batch of datajava.lang.IllegalArgumentException - if the batch engine does not match the trainer enginepublic void step()
public Metrics getMetrics()
public void setMetrics(Metrics metrics)
metrics - the Metrics classpublic java.util.List<Device> getDevices()
public void endEpoch()
public Loss getLoss()
Loss function of the trainer.Loss functionpublic Model getModel()
public java.util.List<Evaluator> getEvaluators()
Evaluators.public TrainingResult getTrainingResult()
TrainingResult.TrainingResultprotected void finalize()
throws java.lang.Throwable
finalize in class java.lang.Objectjava.lang.Throwablepublic void close()
close in interface java.lang.AutoCloseable