Contents

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

engine

Engine executing screen kernels and holding count buffer.

factor

Single-element float tensor scaled for backward seeding.

parameter_count

Number of parameter gradients screened.

options

Configuration options for scale adjustments.

LossScaler(const LossScaler &)=delete
LossScaler & operator=(const LossScaler &)=delete
std::size_t record_screens(const std::vector< SushiBLAS::Tensor * > &gradients)

Records non-finite gradient count kernels for each parameter.

Parameters

gradients

Pointers to gradient tensors in slot order.

Returns

Number of device tasks recorded into the execution plan.

Exceptions

Error

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

Returns pointer to device non-finite count buffer for gating.

Returns

Pointer to base element of count buffer.

std::size_t parameter_count() const noexcept

How many gradients this scaler screens.

double scale() const noexcept

The scale the next step will run at.

double inverse_scale() const noexcept

1 / scale, exactly; what the update multiplies the gradient by.

bool last_step_skipped() const noexcept

Whether the step just completed was thrown away.

std::uint64_t skipped_steps() const noexcept

How many steps have been thrown away in total.

std::uint64_t backoffs() const noexcept

How many times the scale has halved.

std::uint64_t growths() const noexcept

How many times the scale has doubled.

const LossScaleState & state() const noexcept

The state machine, for a test or a report that wants the detail.

std::string to_string() const

A one-line summary for a run report.