class
SushiAI::Optim::Optimizer
Common base class providing parameter update recording and lifecycle hooks.
- Declared in
include/SushiAI/optim/optimizer.hpp- Inherits
SushiAI::Graph::IStepRecorder
Public member functions
Optimizer(SushiBLAS::Engine &engine, std::vector< ParameterSlot > parameters)Binds an optimizer to a fixed set of parameters.
Parameters
engineThe engine the update is recorded on; must outlive this object.
parametersThe parameters to update, with their gradients.
Exceptions
ErrorIf a slot's parameter and gradient disagree in shape or dtype, or the element type is not a real float.
~Optimizer() override=defaultOptimizer(const Optimizer &)=deleteOptimizer & operator=(const Optimizer &)=deletevoid step()Applies one update step standalone to every registered parameter.
Exceptions
ErrorIf optimizer has already been recorded into a compiled step.
virtual bool is_replay_invariant() const noexcept=0Indicates whether recorded update operations survive graph replays unchanged.
Returns
True if updates can be embedded directly into compiled training steps.
std::size_t parameter_count() const noexceptHow many parameters this optimizer updates.
std::uint64_t step_count() const noexceptHow many updates have been applied; 0 before the first step().
const std::vector< ParameterSlot > & parameters() const noexceptThe parameters, for inspection.
void set_loss_scaler(LossScaler *scaler)Attaches a gradient loss scaler to this optimizer.
Parameters
scalerLoss scaler to attach, or nullptr to detach.
Exceptions
ErrorIf optimizer updates are not replay-invariant.
LossScaler * loss_scaler() const noexceptThe scaler attached to this optimizer, or null.
virtual void before_replay() finalRefreshes this step's numbers into device memory.
virtual void after_replay() finalAdvances loss scale state machine followed by step counter.
virtual std::size_t record_into_step() finalRecords parameter update operations into compiled execution plan.
Returns
Number of device tasks recorded.
bool is_in_compiled_step() const noexceptWhether a compiled step has taken ownership of this update.
True from the moment a Graph::CompiledStep recorded it.
Protected member functions
virtual std::size_t record_update()=0Records one update for every parameter; the subclass's job.
Records only - nothing is executed and nothing is waited on.
Returns
How many device tasks were recorded.
std::size_t record_all()Records loss scaler screens followed by parameter updates.
Returns
Total number of device tasks recorded.
virtual void write_step_scalars()=0Writes current step scalars into device memory before execution.
float inverse_loss_scale() const noexceptReturns multiplier for unscaling gradients before updates.
Returns
Reciprocal of current loss scale factor.
const std::int64_t * gradient_gate() const noexceptReturns pointer to device non-finite count buffer, or nullptr.
Returns
Pointer to per-parameter non-finite gradient counts.
void flush()Runs the recorded update and waits for the device.
SushiBLAS::Tensor allocate_state(const SushiBLAS::Tensor &like) constAllocates a state buffer shaped like like, with its own Storage.
SushiBLAS::Engine & engine() const noexceptThe engine the update is recorded on.
std::vector< ParameterSlot > & slots() noexceptThe parameters being updated.
void advance() noexceptBumps the step counter; call once per step().

