Contents

class

SushiAI::Optim::SGD

The update rule above, applied to a fixed parameter set.

Declared in
include/SushiAI/optim/sgd.hpp
Inherits
SushiAI::Optim::Optimizer

Public member functions

SGD(SushiBLAS::Engine &engine, std::vector< ParameterSlot > parameters, SGDOptions options={})

Binds SGD to parameters and allocates momentum state buffers.

Parameters

engine

Engine used to record update kernels.

parameters

Parameter and gradient slots to update.

options

Hyperparameters configuring learning rate and momentum.

Exceptions

Error

If nesterov is set without momentum, or with dampening.

virtual bool is_replay_invariant() const noexcept override

True when the update is fused, and therefore replayable.

const SGDOptions & options() const noexcept

The hyperparameters in force.

SGDOptions & options() noexcept

Mutable access, for a learning-rate schedule.

Protected member functions

virtual std::size_t record_update() override

Records the update for every parameter into the engine's builder.

Returns

One task per parameter when fused, four or five each when not.

virtual void write_step_scalars() override

Recomputes this step's scalars into the device buffer.

Exceptions

Error

If momentum is enabled without allocated buffer.