//! Reduction modes for loss functions /// Reduction mode for aggregating loss values #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Reduction { /// No reduction - return per-sample/per-element losses None, /// Mean reduction - return average loss Mean, /// Sum reduction - return sum of losses Sum, } impl Default for Reduction { fn default() -> Self { Self::Mean } } impl Reduction { /// Apply the reduction to a tensor of losses pub fn apply(&self, losses: &rtx_tensor::Tensor) -> crate::Result { match self { Reduction::None => Ok(losses.clone()), Reduction::Mean => { // For mean, first sum all elements, then divide by total count let sum_result = losses.sum(None)?; let total_elements = losses.numel() as f32; sum_result.div_scalar(total_elements).map_err(Into::into) }, Reduction::Sum => { // Sum all elements losses.sum(None).map_err(Into::into) }, } } /// Apply reduction over specific dimensions pub fn apply_dims(&self, losses: &rtx_tensor::Tensor, dim: Option) -> crate::Result { match self { Reduction::None => Ok(losses.clone()), Reduction::Mean => { // For mean, first sum along the dimension, then divide by the size of that dimension let sum_result = losses.sum(dim)?; if let Some(d) = dim { let dim_size = losses.shape().dims()[d] as f32; sum_result.div_scalar(dim_size).map_err(Into::into) } else { let total_elements = losses.numel() as f32; sum_result.div_scalar(total_elements).map_err(Into::into) } }, Reduction::Sum => losses.sum(dim).map_err(Into::into), } } }