// cross_entropy.mbt 鈥?cross-entropy loss (v0.13.2).
//
// Given:
// - `log_probs`: Tensor of shape [batch, n_classes], the
// log-softmax output of a network.
// - `targets`: Tensor of shape [batch], integer class indices
// (each value in [0, n_classes)).
//
// Returns a Tensor of shape [1] containing the **mean** loss:
// loss = -mean over batch of log_probs[batch, targets[batch]]
//
// Also returns the per-batch loss as a Tensor of shape [batch].
// We bundle both into a struct so callers can choose.
///|
/// Cross-entropy loss result. `mean` is a 1-element Tensor
/// (`shape=[1]`) and `per_batch` has shape `[batch]`.
pub struct CrossEntropyLoss {
mean : Tensor
per_batch : Tensor
}
///|
/// Compute the cross-entropy loss given log-probabilities and
/// integer target indices.
pub fn cross_entropy_loss(
log_probs : Tensor,
targets : Tensor,
) -> CrossEntropyLoss {
if log_probs.shape.length() != 2 {
abort("cross_entropy_loss: log_probs must be 2D [batch, n_classes]")
}
if targets.shape.length() != 1 {
abort("cross_entropy_loss: targets must be 1D [batch]")
}
let batch = log_probs.shape[0]
let n_classes = log_probs.shape[1]
let per_batch : Array[Float] = Array::make(batch, 0.0F)
if targets.data.length() != batch {
abort("cross_entropy_loss: targets length mismatch")
}
let mut total : Float = 0.0F
for b in 0.. Int.
// (This assumes the user stored integer-valued floats.)
let t_int = t.to_int()
if t_int < 0 || t_int >= n_classes {
abort("cross_entropy_loss: target index out of range")
}
let lp = log_probs.data[b * n_classes + t_int]
let loss = -lp
per_batch[b] = loss
total = total + loss
}
let mean = Tensor::from([total / Float::from_int(batch)], [1])
{ mean, per_batch: Tensor::from(per_batch, [batch]) }
}