use crate::core::autograd::BackwardContext;
use crate::tensor::Tensor;
use ndarray::Axis;
use std::rc::Rc;
pub fn sparse_cross_entropy_op(logits: &Tensor, targets: &Tensor) -> Tensor {
let logits_data = logits.data.borrow();
let targets_data = targets.data.borrow();
let last_axis = Axis(logits_data.ndim() - 1);
let max_logits = logits_data
.map_axis(last_axis, |row| {
row.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b))
})
.into_dyn()
.insert_axis(last_axis);
let stable_logits = &*logits_data - &max_logits;
let log_sum_exp = stable_logits
.mapv(|v| v.exp())
.sum_axis(last_axis)
.mapv(|v| v.ln())
.into_dyn()
.insert_axis(last_axis);
let log_probs = &stable_logits - &log_sum_exp;
let seq_len = targets_data.len();
let mut total_loss = 0.0;
for i in 0..seq_len {
let target_idx = targets_data[i] as usize;
total_loss -= log_probs[[0, i, target_idx]];
}
let mean_loss = total_loss / seq_len as f32;
let mut result = Tensor::new(ndarray::arr0(mean_loss).into_dyn(), logits.grad.is_some());
if logits.grad.is_some() {
let logits_for_closure = logits.clone();
let logits_for_inputs = logits.clone();
let targets_for_closure = targets.clone();
let probabilities = log_probs.mapv(|v| v.exp());
let backward_fn = Box::new(move |upstream_grad: &ndarray::ArrayD<f32>| {
if let Some(logits_grad) = &logits_for_closure.grad {
let upstream_scalar = *upstream_grad.first().expect("Upstream grad for loss must be a scalar");
let seq_len_bw = targets_for_closure.data.borrow().len();
let mut d_logits = probabilities.clone();
let targets_data_bw = targets_for_closure.data.borrow();
for i in 0..seq_len_bw {
let target_idx = targets_data_bw[i] as usize;
d_logits[[0, i, target_idx]] -= 1.0;
}
d_logits.mapv_inplace(|v| v * upstream_scalar / seq_len_bw as f32);
logits_grad.borrow_mut().scaled_add(1.0, &d_logits);
}
});
result.ctx = Some(Rc::new(BackwardContext {
inputs: vec![logits_for_inputs],
backward_fn,
}));
}
result
}