Contents

namespace

SushiAI::Train

Declared in
include/SushiAI/train/objective.hpp

Contains

Enumerations

enum class OptimizerPlacement

Whether the parameter update is part of the compiled step.

IN_PLAN

Recorded into the same plan as forward and backward; the default.

PER_STEP

Stepped on its own after every replay, with its own flush.

Typedefs

using SushiAI::Train::ObjectiveFactory = std::unique_ptr<IObjective> (*)(int64_t batch, int64_t classes, const Config::Value& attributes, std::string_view context)

Builds one objective from its configuration attributes.

Parameters

batch

Samples per step.

classes

How many classes the dataset spans.

attributes

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

context

What to call this objective in a diagnostic.

Returns

The objective.

Variables

constexpr std::array< ObjectiveEntry, 1 > OBJECTIVE_REGISTRY = { ObjectiveEntry{"cross_entropy", &make_cross_entropy}, }

Every objective a configuration file may name.

A file-scope constant table; nothing self-registers. See config/registry.hpp for why that is forced rather than preferred.

Functions

std::unique_ptr< IObjective > make_cross_entropy(int64_t batch, int64_t classes, const Config::Value &attributes, std::string_view context)

Builds a cross-entropy objective from attributes.

Parameters

batch

Samples per step.

classes

Number of target classes.

attributes

Objective attribute values.

context

Diagnostic context string.

Returns

The constructed objective instance.

std::unique_ptr< IObjective > make_objective(std::string_view type, int64_t batch, int64_t classes, const Config::Value &attributes, std::string_view context)

Builds the objective named by type.

Parameters

type

Registry key; batch Samples per step.

classes

Target classes; attributes Configuration.

context

Diagnostic context string.

Returns

The constructed objective instance.

template <typename M>
TrainingGraph build_graph(M &module, const IObjective &objective, int64_t batch, int64_t features, DType dtype=DType::FLOAT32, Graph::FusionOptions fusion={}, const Graph::IPrecisionPolicy *precision=nullptr, bool loss_scaling=false)

Traces module into a full training graph and differentiates it.

Parameters

module

The model to trace.

objective

Training loss and metric objective.

batch

Samples per step.

features

Input features per sample.

Returns

The constructed training graph.

template <typename M>
TrainingGraph build_classifier(M &module, int64_t batch, int64_t features, int64_t classes, DType dtype=DType::FLOAT32, Graph::FusionOptions fusion={}, const Graph::IPrecisionPolicy *precision=nullptr, bool loss_scaling=false)

Traces module into a cross-entropy classifier training graph.

Parameters

module

The model to trace.

batch

Samples per step.

features

Input features per sample.

classes

Target class count.

Returns

The constructed training graph.

std::vector< Optim::ParameterSlot > bind_parameters(const TrainingGraph &built, Graph::TensorPool &pool, Graph::PlanOptions options={})

Plans memory, allocates leaves, and pairs parameters with gradients.

Parameters

built

The traced graph.

pool

Destination tensor pool.

options

Memory packing options.

Returns

Parameter slots paired with gradients.