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) noexceptConstructs the objective for a fixed batch and class count.
Parameters
batchSamples per step.
classesHow many classes.
dtypeThe element type the target leaf carries.
virtual void declare(const Graph::GraphBuilder &builder, ObjectiveLeaves &leaves) const overrideDeclares whatever leaves the graph needs beyond the model's input.
Parameters
builderWhere the leaves go.
leavesWhere to record them, by role.
virtual Graph::ValueId loss(const Graph::GraphBuilder &builder, const Graph::Graph &graph, Graph::ValueId output, const ObjectiveLeaves &leaves) const overrideComputes scalar loss from model output.
Parameters
builderWhere the nodes go.
graphThe graph, for checking shapes.
outputThe model output value.
leavesDeclared auxiliary leaves.
Returns
The scalar loss value identifier.
virtual std::optional< Metric > score(const Data::Batch &batch, const SushiBLAS::Tensor &output) const overrideAn optional per-step metric: accuracy, mAP, an L2 residual, or none.
Parameters
batchThe batch that was just stepped.
outputThe model's output tensor.
Returns
The metric, or std::nullopt when this objective offers none.
virtual std::string_view name() const noexcept overrideA short name for a run report.

