Training
Training SameDiff models — TrainingConfig, fit(), listeners, loss curves, and evaluation
Overview of the Training Flow
TrainingConfig
import org.nd4j.autodiff.samediff.TrainingConfig;
import org.nd4j.linalg.learning.config.Adam;
import org.nd4j.linalg.api.buffer.DataType;
TrainingConfig config = TrainingConfig.builder()
.updater(new Adam(1e-3)) // optimizer with learning rate
.dataSetFeatureMapping("input") // feature array -> placeholder
.dataSetLabelMapping("labels") // label array -> placeholder
.lossVariables("loss") // which SDVariable holds the scalar loss
.build();
sd.setTrainingConfig(config);Required settings
Setting
Method
Notes
Optimizer (IUpdater)
MultiDataSet mappings
Data type conversion
fit() — Running Training
With a DataSetIterator
DataSetIteratorWith a MultiDataSetIterator
MultiDataSetIteratorWith a validation set
Fitting a single batch manually
History and LossCurve
Listeners
Available built-in listeners
ScoreIterationListener
SameDiffListener interface
CheckpointListener
Adding Evaluation During Training
End-to-End Training Example
Controlling Which Parameters Are Trained
Gradient Clipping
Last updated
Was this helpful?