Contents

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

engine

The engine the update is recorded on; must outlive this object.

parameters

The parameters to update, with their gradients.

Exceptions

Error

If a slot's parameter and gradient disagree in shape or dtype, or the element type is not a real float.

~Optimizer() override=default
Optimizer(const Optimizer &)=delete
Optimizer & operator=(const Optimizer &)=delete
void step()

Applies one update step standalone to every registered parameter.

Exceptions

Error

If optimizer has already been recorded into a compiled step.

virtual bool is_replay_invariant() const noexcept=0

Indicates 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 noexcept

How many parameters this optimizer updates.

std::uint64_t step_count() const noexcept

How many updates have been applied; 0 before the first step().

const std::vector< ParameterSlot > & parameters() const noexcept

The parameters, for inspection.

void set_loss_scaler(LossScaler *scaler)

Attaches a gradient loss scaler to this optimizer.

Parameters

scaler

Loss scaler to attach, or nullptr to detach.

Exceptions

Error

If optimizer updates are not replay-invariant.

LossScaler * loss_scaler() const noexcept

The scaler attached to this optimizer, or null.

virtual void before_replay() final

Refreshes this step's numbers into device memory.

virtual void after_replay() final

Advances loss scale state machine followed by step counter.

virtual std::size_t record_into_step() final

Records parameter update operations into compiled execution plan.

Returns

Number of device tasks recorded.

bool is_in_compiled_step() const noexcept

Whether 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()=0

Records 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()=0

Writes current step scalars into device memory before execution.

float inverse_loss_scale() const noexcept

Returns multiplier for unscaling gradients before updates.

Returns

Reciprocal of current loss scale factor.

const std::int64_t * gradient_gate() const noexcept

Returns 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) const

Allocates a state buffer shaped like like, with its own Storage.

SushiBLAS::Engine & engine() const noexcept

The engine the update is recorded on.

std::vector< ParameterSlot > & slots() noexcept

The parameters being updated.

void advance() noexcept

Bumps the step counter; call once per step().