struct
SushiAI::NN::CrossEntropyLoss
Computes mean cross-entropy over a batch from raw logits.
- Declared in
include/SushiAI/nn/loss.hpp
Public member functions
Graph::ValueId forward(const Graph::GraphBuilder &builder, Graph::ValueId logits, Graph::ValueId one_hot) constEmits the loss for one batch.
Parameters
builderWhere the nodes go.
logitsAn [N, C] value of unnormalised scores.
one_hotAn [N, C] one-hot target matrix.
Returns
A 1-element loss value's id.
Exceptions
ErrorIf the two operands disagree in shape or dtype.

