use crate::compat::*;
use crate::module::Module;
use hodu_core::{error::HoduResult, scalar::Scalar, tensor::Tensor};
#[derive(Module, Clone)]
#[module(inputs = 2)]
pub struct BCELoss {
epsilon: Scalar,
}
impl BCELoss {
pub fn new() -> Self {
Self {
epsilon: Scalar::F32(1e-7), }
}
pub fn with_epsilon(epsilon: impl Into<Scalar>) -> Self {
Self {
epsilon: epsilon.into(),
}
}
pub fn forward(&self, (pred, target): (&Tensor, &Tensor)) -> HoduResult<Tensor> {
let one_minus_eps_scalar = Scalar::one(self.epsilon.get_dtype()) - self.epsilon;
let pred_clamped = pred.clamp(self.epsilon, one_minus_eps_scalar)?;
let log_pred = pred_clamped.ln()?;
let first_term = target.mul(&log_pred)?;
let one = Tensor::ones_like(target)?;
let one_minus_target = one.sub(target)?;
let one_minus_pred = Tensor::ones_like(&pred_clamped)?.sub(&pred_clamped)?;
let log_one_minus_pred = one_minus_pred.ln()?;
let second_term = one_minus_target.mul(&log_one_minus_pred)?;
let bce = first_term.add(&second_term)?.neg()?;
bce.mean_all()
}
}
#[derive(Module, Clone)]
#[module(inputs = 2)]
pub struct BCEWithLogitsLoss;
impl BCEWithLogitsLoss {
pub fn new() -> Self {
Self
}
pub fn forward(&self, (logits, target): (&Tensor, &Tensor)) -> HoduResult<Tensor> {
let zeros = Tensor::zeros_like(logits)?;
let max_val = logits.maximum(&zeros)?;
let neg_abs = logits.abs()?.neg()?; let log_term = neg_abs.exp()?.add_scalar(Scalar::one(neg_abs.get_dtype()))?.ln()?;
let target_term = logits.mul(target)?;
let loss = max_val.sub(&target_term)?.add(&log_term)?;
loss.mean_all()
}
}