Files
rustytorch/crates/core/rtx-losses/src/reduction.rs
T
2026-03-04 00:08:42 +00:00

56 lines
2.0 KiB
Rust

//! 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<rtx_tensor::Tensor> {
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<usize>) -> crate::Result<rtx_tensor::Tensor> {
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),
}
}
}