56 lines
2.0 KiB
Rust
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),
|
|
}
|
|
}
|
|
} |