Contents

namespace

SushiAI::Optim

Declared in
include/SushiAI/optim/adamw.hpp

Contains

Enumerations

enum class UpdateKernel

How an optimizer spells its per-parameter update on the device.

FUSED

One fused kernel per parameter; the default.

SEPARATE

The original sequence of separate SushiBLAS calls.

Typedefs

using SushiAI::Optim::OptimizerFactory = std::unique_ptr<Optimizer> (*)( SushiBLAS::Engine& engine, std::vector<ParameterSlot> parameters, const Config::Value& attributes, std::string_view context)

Builds one optimizer over parameters from its configuration.

Parameters

engine

The engine its state buffers are allocated from.

parameters

The parameter/gradient pairs to update; moved in.

attributes

The optimizer's own keys; the factory rejects any it does not know.

context

What to call this optimizer in a diagnostic.

Returns

The optimizer.

Variables

constexpr std::array< OptimizerEntry, 2 > OPTIMIZER_REGISTRY = { OptimizerEntry{"adamw", &make_adamw}, OptimizerEntry{"sgd", &make_sgd}, }

Every optimizer a configuration file may name.

Functions

std::unique_ptr< Optimizer > make_adamw(SushiBLAS::Engine &engine, std::vector< ParameterSlot > parameters, const Config::Value &attributes, std::string_view context)

Builds an AdamW optimizer from configuration attributes.

Parameters

engine

Engine used to allocate optimizer moment buffers.

parameters

Parameter and gradient slots moved into optimizer.

attributes

Optional hyperparameters (learning_rate, beta1, beta2, etc.).

context

Descriptive string for diagnostic error messages.

Returns

Constructed AdamW optimizer instance.

std::unique_ptr< Optimizer > make_sgd(SushiBLAS::Engine &engine, std::vector< ParameterSlot > parameters, const Config::Value &attributes, std::string_view context)

Builds an SGD optimizer from configuration attributes.

Parameters

engine

Engine used to allocate optimizer momentum buffers.

parameters

Parameter and gradient slots moved into optimizer.

attributes

Optional hyperparameters (learning_rate, momentum, etc.).

context

Descriptive string for diagnostic error messages.

Returns

Constructed SGD optimizer instance.

std::unique_ptr< Optimizer > make_optimizer(SushiBLAS::Engine &engine, std::string_view type, std::vector< ParameterSlot > parameters, const Config::Value &attributes, std::string_view context)

Builds the optimizer corresponding to the specified type key.

Parameters

engine

Engine used to allocate optimizer state.

type

Registry key identifying optimizer type.

parameters

Parameter and gradient slots moved into optimizer.

attributes

Optimizer configuration attributes.

Returns

Constructed optimizer instance.