use burn_core as burn;
use burn::config::Config;
use burn::module::{Content, DisplaySettings, Module, ModuleDisplay};
use burn::tensor::Tensor;
use super::Reduction;
const HALF_LN_TWO_PI: f64 = 0.918_938_533_204_672_8;
#[derive(Config, Debug)]
pub struct GaussianNLLLossConfig {
#[config(default = 1e-6)]
pub eps: f64,
#[config(default = false)]
pub full: bool,
}
impl GaussianNLLLossConfig {
pub fn init(&self) -> GaussianNLLLoss {
GaussianNLLLoss {
eps: self.eps,
full: self.full,
}
}
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct GaussianNLLLoss {
pub eps: f64,
pub full: bool,
}
impl ModuleDisplay for GaussianNLLLoss {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
content
.add("eps", &self.eps)
.add("full", &self.full)
.optional()
}
}
impl GaussianNLLLoss {
pub fn forward<const D: usize>(
&self,
input: Tensor<D>,
target: Tensor<D>,
var: Tensor<D>,
reduction: Reduction,
) -> Tensor<1> {
let loss = self.forward_no_reduction(input, target, var);
match reduction {
Reduction::Mean | Reduction::Auto => loss.mean(),
Reduction::Sum => loss.sum(),
other => panic!("{other:?} reduction is not supported"),
}
}
pub fn forward_no_reduction<const D: usize>(
&self,
input: Tensor<D>,
target: Tensor<D>,
var: Tensor<D>,
) -> Tensor<D> {
let var = var.clamp_min(self.eps);
let squared_error = (input - target).square();
let loss = var
.clone()
.log()
.add(squared_error.div(var))
.mul_scalar(0.5);
if self.full {
loss.add_scalar(HALF_LN_TWO_PI)
} else {
loss
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::TensorData;
use burn::tensor::Tolerance;
type FT = f32;
#[test]
fn test_gaussian_nll_loss() {
let device = Default::default();
let input = Tensor::<1>::from_data(TensorData::from([1.0, 2.0]), &device);
let target = Tensor::<1>::from_data(TensorData::from([1.5, 1.0]), &device);
let var = Tensor::<1>::from_data(TensorData::from([0.5, 1.0]), &device);
let loss = GaussianNLLLossConfig::new().init();
let no_reduction = loss.forward_no_reduction(input.clone(), target.clone(), var.clone());
let mean = loss.forward(input.clone(), target.clone(), var.clone(), Reduction::Mean);
let sum = loss.forward(input, target, var, Reduction::Sum);
let expected = TensorData::from([-0.096574, 0.5]);
no_reduction
.into_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
mean.into_data()
.assert_approx_eq::<FT>(&TensorData::from([0.201713]), Tolerance::default());
sum.into_data()
.assert_approx_eq::<FT>(&TensorData::from([0.403426]), Tolerance::default());
}
#[test]
fn test_gaussian_nll_loss_full() {
let device = Default::default();
let input = Tensor::<1>::from_data(TensorData::from([1.0, 2.0]), &device);
let target = Tensor::<1>::from_data(TensorData::from([1.5, 1.0]), &device);
let var = Tensor::<1>::from_data(TensorData::from([0.5, 1.0]), &device);
let loss = GaussianNLLLossConfig::new().with_full(true).init();
let no_reduction = loss.forward_no_reduction(input, target, var);
let expected = TensorData::from([0.822365, 1.418939]);
no_reduction
.into_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
}
#[test]
fn test_gaussian_nll_loss_clamps_var_to_eps() {
let device = Default::default();
let input = Tensor::<1>::from_data(TensorData::from([1.0, 2.0]), &device);
let target = Tensor::<1>::from_data(TensorData::from([1.5, 1.0]), &device);
let loss = GaussianNLLLossConfig::new().init();
let var_below = Tensor::<1>::from_data(TensorData::from([1e-8, 1e-8]), &device);
let var_eps = Tensor::<1>::from_data(TensorData::from([1e-6, 1e-6]), &device);
let clamped = loss.forward_no_reduction(input.clone(), target.clone(), var_below);
let at_eps = loss.forward_no_reduction(input, target, var_eps);
clamped
.into_data()
.assert_approx_eq::<FT>(&at_eps.into_data(), Tolerance::default());
}
#[test]
fn display() {
let config = GaussianNLLLossConfig::new().with_eps(0.5);
let loss = config.init();
assert_eq!(
alloc::format!("{loss}"),
"GaussianNLLLoss {eps: 0.5, full: false}"
);
}
}