Contents

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

engine

Execution engine; graph Traced graph.

pool

Tensor pool; spec Step boundary values.

optimizer

Parameter optimizer; options Configuration.

Trainer(const Trainer &)=delete
Trainer & operator=(const Trainer &)=delete
StepMetrics run_step(const Data::Batch &batch)

Runs one step: feed, replay, read metrics, update.

Parameters

batch

The batch to train on.

Returns

The step's loss and accuracy.

Exceptions

Error

If batch does not match the compiled shapes.

RunReport fit(const Data::Dataset &dataset)

Runs options.epochs passes over dataset.

Parameters

dataset

The batches to train on.

Returns

The run's summary.

Exceptions

Error

If the dataset's shapes do not match the compiled step.

std::size_t compile_count() const noexcept

How many times the step has been compiled; one, always.

std::size_t step_count() const noexcept

How many steps have run.

bool optimizer_is_in_plan() const noexcept

Whether the optimizer's update is part of the compiled step.

const Graph::CompiledStep & step() const noexcept

The compiled step, for a test that wants to inspect it.

double current_loss() const

Reads the loss tensor's current value from the host.

TrainerOptions & options() noexcept

The options in force.