class
SushiAI::Optim::AdamW
The update above, applied to a fixed parameter set.
- Declared in
include/SushiAI/optim/adamw.hpp- Inherits
SushiAI::Optim::Optimizer
Public member functions
AdamW(SushiBLAS::Engine &engine, std::vector< ParameterSlot > parameters, AdamWOptions options={})Binds AdamW to parameters, allocates and zeroes its state.
Parameters
engineThe engine the update is recorded on.
parametersThe parameters to update.
optionsThe hyperparameters.
Exceptions
ErrorIf a beta is outside [0, 1) or eps is not positive.
virtual bool is_replay_invariant() const noexcept overrideTrue when the update is fused, and therefore replayable.
const AdamWOptions & options() const noexceptThe hyperparameters in force.
AdamWOptions & options() noexceptMutable access, for a learning-rate schedule.
const SushiBLAS::Tensor & first_moment(std::size_t index) constReturns parameter's first moment tensor for inspection.
Parameters
indexParameter slot index in registration order.
Returns
First moment tensor reference.
Exceptions
ErrorIf index is out of range.
const SushiBLAS::Tensor & second_moment(std::size_t index) constParameter index's second moment, v.
Parameters
indexWhich parameter, in slot order.
Returns
Its second moment.
Exceptions
ErrorIf
indexis out of range.
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; twelve with decay and eleven without, when not.
virtual void write_step_scalars() overrideRecomputes this step's scalars into the device buffer.

