use crate::error::Error;
use crate::neural_network::Tensor;
use crate::neural_network::losses::{
normalize_and_clip_rows, stable_log_softmax_softmax, validate_same_shape,
};
use crate::neural_network::traits::Loss;
use ndarray::{Array2, Zip};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct CategoricalCrossEntropy {
from_logits: bool,
}
impl CategoricalCrossEntropy {
pub fn new(from_logits: bool) -> Self {
Self { from_logits }
}
}
fn leading_and_classes(t: &Tensor) -> (usize, usize) {
let ndim = t.ndim();
let classes = t.shape()[ndim - 1];
let sites: usize = t.shape()[..ndim - 1].iter().product();
(sites, classes)
}
fn validate_shapes(y_true: &Tensor, y_pred: &Tensor) -> Result<(), Error> {
if y_true.is_empty() {
return Err(Error::empty_input(
"CategoricalCrossEntropy expects non-empty y_true",
));
}
if y_true.ndim() < 2 {
return Err(Error::invalid_input(format!(
"CategoricalCrossEntropy expects at least 2D tensors [batch, classes], got {}D",
y_true.ndim()
)));
}
validate_same_shape(y_true, y_pred)
}
impl Loss for CategoricalCrossEntropy {
fn compute_loss(&self, y_true: &Tensor, y_pred: &Tensor) -> Result<f32, Error> {
validate_shapes(y_true, y_pred)?;
let (sites, classes) = leading_and_classes(y_pred);
let n = sites as f32;
if self.from_logits {
let logits = y_pred
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE logits reshape failed: {e}")))?;
let labels = y_true
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE labels reshape failed: {e}")))?;
let (log_sm, _) = stable_log_softmax_softmax(&logits.view());
return Ok(-(&labels * &log_sm).sum() / n);
}
let probs = y_pred
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE probability reshape failed: {e}")))?;
let labels = y_true
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE labels reshape failed: {e}")))?;
let (normalized, _) = normalize_and_clip_rows(&probs.view());
Ok(-(&labels * &normalized.mapv(f32::ln)).sum() / n)
}
fn compute_grad(&self, y_true: &Tensor, y_pred: &Tensor) -> Result<Tensor, Error> {
validate_shapes(y_true, y_pred)?;
let (sites, classes) = leading_and_classes(y_pred);
let n = sites as f32;
if self.from_logits {
let logits = y_pred
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE logits reshape failed: {e}")))?;
let labels = y_true
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE labels reshape failed: {e}")))?;
let (_, sm) = stable_log_softmax_softmax(&logits.view());
let grad2d = (&sm - &labels) / n;
let grad = grad2d
.into_shape_with_order(y_pred.raw_dim())
.map_err(|e| Error::computation(format!("CCE gradient reshape failed: {e}")))?;
return Ok(grad);
}
let probs = y_pred
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE probability reshape failed: {e}")))?;
let labels = y_true
.to_shape((sites, classes))
.map_err(|e| Error::computation(format!("CCE labels reshape failed: {e}")))?;
let (normalized, row_sums) = normalize_and_clip_rows(&probs.view());
let mut grad2d = Array2::<f32>::zeros((sites, classes));
Zip::from(grad2d.rows_mut())
.and(normalized.rows())
.and(labels.rows())
.and(&row_sums)
.for_each(|mut grad_row, normalized_row, label_row, &sum| {
let target_mass = label_row.sum();
Zip::from(&mut grad_row)
.and(normalized_row)
.and(label_row)
.for_each(|grad, &q, &y| *grad = (target_mass - y / q) / (n * sum));
});
grad2d
.into_shape_with_order(y_pred.raw_dim())
.map_err(|e| Error::computation(format!("CCE gradient reshape failed: {e}")))
}
}