Contents

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

engine

The engine the update is recorded on.

parameters

The parameters to update.

options

The hyperparameters.

Exceptions

Error

If a beta is outside [0, 1) or eps is not positive.

virtual bool is_replay_invariant() const noexcept override

True when the update is fused, and therefore replayable.

const AdamWOptions & options() const noexcept

The hyperparameters in force.

AdamWOptions & options() noexcept

Mutable access, for a learning-rate schedule.

const SushiBLAS::Tensor & first_moment(std::size_t index) const

Returns parameter's first moment tensor for inspection.

Parameters

index

Parameter slot index in registration order.

Returns

First moment tensor reference.

Exceptions

Error

If index is out of range.

const SushiBLAS::Tensor & second_moment(std::size_t index) const

Parameter index's second moment, v.

Parameters

index

Which parameter, in slot order.

Returns

Its second moment.

Exceptions

Error

If index is out of range.

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; twelve with decay and eleven without, when not.

virtual void write_step_scalars() override

Recomputes this step's scalars into the device buffer.