class
SushiAI::Train::Trainer
Owns the compiled step and drives it over a dataset.
- Declared in
include/SushiAI/train/trainer.hpp
Non-copyable and non-movable, because it holds a CompiledStep, which holds raw node pointers into its own storage.
Public member functions
Trainer(SushiBLAS::Engine &engine, const Graph::Graph &graph, Graph::TensorPool &pool, StepSpec spec, Optim::Optimizer &optimizer, TrainerOptions options={})Lowers and compiles the execution step.
Parameters
engineExecution engine;
graphTraced graph.poolTensor pool;
specStep boundary values.optimizerParameter optimizer;
optionsConfiguration.
Trainer(const Trainer &)=deleteTrainer & operator=(const Trainer &)=deleteStepMetrics run_step(const Data::Batch &batch)Runs one step: feed, replay, read metrics, update.
Parameters
batchThe batch to train on.
Returns
The step's loss and accuracy.
Exceptions
ErrorIf
batchdoes not match the compiled shapes.
RunReport fit(const Data::Dataset &dataset)Runs options.epochs passes over dataset.
Parameters
datasetThe batches to train on.
Returns
The run's summary.
Exceptions
ErrorIf the dataset's shapes do not match the compiled step.
std::size_t compile_count() const noexceptHow many times the step has been compiled; one, always.
std::size_t step_count() const noexceptHow many steps have run.
bool optimizer_is_in_plan() const noexceptWhether the optimizer's update is part of the compiled step.
const Graph::CompiledStep & step() const noexceptThe compiled step, for a test that wants to inspect it.
double current_loss() constReads the loss tensor's current value from the host.
TrainerOptions & options() noexceptThe options in force.

