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
engineEngine used to record update kernels.
parametersParameter and gradient slots to update.
optionsHyperparameters configuring learning rate and momentum.
Exceptions
ErrorIf nesterov is set without momentum, or with dampening.
virtual bool is_replay_invariant() const noexcept overrideTrue when the update is fused, and therefore replayable.
const SGDOptions & options() const noexceptThe hyperparameters in force.
SGDOptions & options() noexceptMutable access, for a learning-rate schedule.
Protected member functions
virtual std::size_t record_update() overrideRecords 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() overrideRecomputes this step's scalars into the device buffer.
Exceptions
ErrorIf momentum is enabled without allocated buffer.

