Contents

class

SushiAI::Train::CrossEntropyObjective

Multi-class cross-entropy over one-hot targets, and accuracy.

Declared in
include/SushiAI/train/objective.hpp
Inherits
SushiAI::Train::IObjective

Emits nll_loss(log_softmax(output), targets) - the same two operations in the same order as the code this replaced, so the existing suite is the regression test.

Public member functions

CrossEntropyObjective(int64_t batch, int64_t classes, DType dtype=DType::FLOAT32) noexcept

Constructs the objective for a fixed batch and class count.

Parameters

batch

Samples per step.

classes

How many classes.

dtype

The element type the target leaf carries.

virtual void declare(const Graph::GraphBuilder &builder, ObjectiveLeaves &leaves) const override

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 override

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 override

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 override

A short name for a run report.