namespace
SushiAI::Train
- Declared in
include/SushiAI/train/objective.hpp
Contains
SushiAI::Train::CrossEntropyObjectiveSushiAI::Train::IObjectiveSushiAI::Train::MetricSushiAI::Train::ObjectiveEntrySushiAI::Train::ObjectiveLeavesSushiAI::Train::ParameterBindingSushiAI::Train::RunReportSushiAI::Train::StepMetricsSushiAI::Train::StepSpecSushiAI::Train::TrainerSushiAI::Train::TrainerOptionsSushiAI::Train::TrainingGraph
Enumerations
enum class OptimizerPlacementWhether the parameter update is part of the compiled step.
IN_PLANRecorded into the same plan as forward and backward; the default.
PER_STEPStepped 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
batchSamples per step.
classesHow many classes the dataset spans.
attributesThe objective's own keys; the factory rejects any it does not know.
contextWhat 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
batchSamples per step.
classesNumber of target classes.
attributesObjective attribute values.
contextDiagnostic 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
typeRegistry key;
batchSamples per step.classesTarget classes;
attributesConfiguration.contextDiagnostic 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
moduleThe model to trace.
objectiveTraining loss and metric objective.
batchSamples per step.
featuresInput 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
moduleThe model to trace.
batchSamples per step.
featuresInput features per sample.
classesTarget 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
builtThe traced graph.
poolDestination tensor pool.
optionsMemory packing options.
Returns
Parameter slots paired with gradients.

