class
SushiAI::Optim::LossScaler
The state machine, its device-resident scale, and the screen.
- Declared in
include/SushiAI/optim/loss_scaler.hpp
Non-copyable: kernels capture the addresses of both its buffers, so the object must be the only thing that owns them.
Public member functions
LossScaler(SushiBLAS::Engine &engine, SushiBLAS::Tensor factor, std::size_t parameter_count, LossScaleOptions options={})Constructs a loss scaler bound to device factor and gradient counts.
Parameters
engineEngine executing screen kernels and holding count buffer.
factorSingle-element float tensor scaled for backward seeding.
parameter_countNumber of parameter gradients screened.
optionsConfiguration options for scale adjustments.
LossScaler(const LossScaler &)=deleteLossScaler & operator=(const LossScaler &)=deletestd::size_t record_screens(const std::vector< SushiBLAS::Tensor * > &gradients)Records non-finite gradient count kernels for each parameter.
Parameters
gradientsPointers to gradient tensors in slot order.
Returns
Number of device tasks recorded into the execution plan.
Exceptions
ErrorIf gradient count does not match parameter count.
void before_replay()Publishes current step scale into device memory before replay.
void after_replay()Reads the screens back and advances the state machine.
One host read of parameter_count 8-byte counts, at the step boundary, after the device has been waited on - so it adds no synchronisation the loop was not already paying.
const std::int64_t * gate() const noexceptReturns pointer to device non-finite count buffer for gating.
Returns
Pointer to base element of count buffer.
std::size_t parameter_count() const noexceptHow many gradients this scaler screens.
double scale() const noexceptThe scale the next step will run at.
double inverse_scale() const noexcept1 / scale, exactly; what the update multiplies the gradient by.
bool last_step_skipped() const noexceptWhether the step just completed was thrown away.
std::uint64_t skipped_steps() const noexceptHow many steps have been thrown away in total.
std::uint64_t backoffs() const noexceptHow many times the scale has halved.
std::uint64_t growths() const noexceptHow many times the scale has doubled.
const LossScaleState & state() const noexceptThe state machine, for a test or a report that wants the detail.
std::string to_string() constA one-line summary for a run report.

