Contents

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) const

Emits the loss for one batch.

Parameters

builder

Where the nodes go.

logits

An [N, C] value of unnormalised scores.

one_hot

An [N, C] one-hot target matrix.

Returns

A 1-element loss value's id.

Exceptions

Error

If the two operands disagree in shape or dtype.