Contents

class

SushiAI::Train::IObjective

What a model is being trained to minimise.

Declared in
include/SushiAI/train/objective.hpp

Public member functions

virtual ~IObjective()=default
virtual void declare(const Graph::GraphBuilder &builder, ObjectiveLeaves &leaves) const =0

Declares whatever leaves the graph needs beyond the model's input.

Parameters

builder

Where the leaves go.

leaves

Where to record them, by role.

virtual Graph::ValueId loss(const Graph::GraphBuilder &builder, const Graph::Graph &graph, Graph::ValueId output, const ObjectiveLeaves &leaves) const =0

Computes scalar loss from model output.

Parameters

builder

Where the nodes go.

graph

The graph, for checking shapes.

output

The model output value.

leaves

Declared auxiliary leaves.

Returns

The scalar loss value identifier.

virtual std::optional< Metric > score(const Data::Batch &batch, const SushiBLAS::Tensor &output) const =0

An optional per-step metric: accuracy, mAP, an L2 residual, or none.

Parameters

batch

The batch that was just stepped.

output

The model's output tensor.

Returns

The metric, or std::nullopt when this objective offers none.

virtual std::string_view name() const noexcept=0

A short name for a run report.