use torsh_core::error::{Result, TorshError};
use torsh_tensor::Tensor;
pub fn cross_entropy(
input: &Tensor,
target: &Tensor<i64>,
weight: Option<&Tensor>,
reduction: &str,
ignore_index: Option<i64>,
) -> Result<Tensor> {
let log_probs = crate::functional::activation::log_softmax(input, Some(-1))?;
let input_shape_binding = input.shape();
let input_dims = input_shape_binding.dims();
let (batch_size, num_classes) = match input_dims {
[batch, classes] => (*batch, *classes),
_ => {
return Err(TorshError::InvalidShape(format!(
"cross_entropy expects a 2-D input [N, C], got shape {input_dims:?}"
)))
}
};
let target_vec = target.to_vec()?;
if target_vec.len() != batch_size {
return Err(TorshError::InvalidShape(format!(
"cross_entropy expects one target per sample: input has {batch_size} rows \
but target has {} entries",
target_vec.len()
)));
}
let mut one_hot_data = vec![0.0f32; batch_size * num_classes];
for (i, &target_idx) in target_vec.iter().enumerate() {
if let Some(ignore_idx) = ignore_index {
if target_idx == ignore_idx {
continue; }
}
if target_idx >= 0 && (target_idx as usize) < num_classes {
if let Some(slot) = one_hot_data.get_mut(i * num_classes + target_idx as usize) {
*slot = 1.0;
}
}
}
let one_hot = Tensor::from_data(one_hot_data, vec![batch_size, num_classes], input.device())?;
let neg_log_likelihood = log_probs.mul_op(&one_hot)?.neg()?;
let loss_per_sample = neg_log_likelihood.sum_dim(&[-1], false)?;
let weighted_loss = if let Some(weights) = weight {
let weight_vec = weights.to_vec()?;
let mut weight_data = vec![1.0f32; batch_size];
for (slot, &target_idx) in weight_data.iter_mut().zip(target_vec.iter()) {
if target_idx >= 0 {
if let Some(&class_weight) = weight_vec.get(target_idx as usize) {
*slot = class_weight;
}
}
}
let weight_tensor = Tensor::from_data(weight_data, vec![batch_size], input.device())?;
loss_per_sample.mul_op(&weight_tensor)?
} else {
loss_per_sample
};
apply_reduction(&weighted_loss, reduction, ignore_index, &target_vec)
}
pub fn binary_cross_entropy(
input: &Tensor,
target: &Tensor,
weight: Option<&Tensor>,
reduction: &str,
) -> Result<Tensor> {
let eps = 1e-7; let eps_tensor = torsh_tensor::creation::full_like(input, eps)?;
let ones = torsh_tensor::creation::ones_like(input)?;
let clamped_input = input.maximum(&eps_tensor)?;
let clamped_input = clamped_input.minimum(&ones.sub(&eps_tensor)?)?;
let log_input = clamped_input.log()?;
let one_minus_input = ones.sub(&clamped_input)?;
let log_one_minus_input = one_minus_input.log()?;
let term1 = target.mul_op(&log_input)?;
let one_minus_target = ones.sub(target)?;
let term2 = one_minus_target.mul_op(&log_one_minus_input)?;
let mut loss = term1.add(&term2)?.neg()?;
if let Some(w) = weight {
loss = loss.mul_op(w)?;
}
apply_reduction(&loss, reduction, None, &[])
}
pub fn binary_cross_entropy_with_logits(
input: &Tensor,
target: &Tensor,
weight: Option<&Tensor>,
reduction: &str,
pos_weight: Option<&Tensor>,
) -> Result<Tensor> {
let zeros = torsh_tensor::creation::zeros_like(input)?;
let ones = torsh_tensor::creation::ones_like(input)?;
let max_term = input.maximum(&zeros)?;
let mult_term = input.mul_op(target)?;
let abs_input = input.abs()?;
let neg_abs_input = abs_input.neg()?;
let exp_term = neg_abs_input.exp()?;
let log_term = ones.add(&exp_term)?.log()?;
let mut loss = max_term.sub(&mult_term)?.add(&log_term)?;
if let Some(pos_w) = pos_weight {
let weighted_target = target.mul_op(pos_w)?;
let pos_term = input.mul_op(&weighted_target)?;
loss = loss.add(&pos_term)?;
}
if let Some(w) = weight {
loss = loss.mul_op(w)?;
}
apply_reduction(&loss, reduction, None, &[])
}
pub fn multi_margin_loss(
input: &Tensor,
target: &Tensor<i64>,
p: i32,
margin: f32,
weight: Option<&Tensor>,
reduction: &str,
) -> Result<Tensor> {
let shape_binding = input.shape();
let input_dims = shape_binding.dims();
let (batch_size, num_classes) = match input_dims {
[batch, classes] => (*batch, *classes),
_ => {
return Err(TorshError::InvalidShape(format!(
"multi_margin_loss expects a 2-D input [N, C], got shape {input_dims:?}"
)))
}
};
if p < 1 {
return Err(TorshError::InvalidArgument(format!(
"multi_margin_loss expects p >= 1 (PyTorch documents p in {{1, 2}}), got {p}"
)));
}
let target_data = target.to_vec()?;
if target_data.len() != batch_size {
return Err(TorshError::InvalidShape(format!(
"multi_margin_loss expects one target per sample: input has {batch_size} rows \
but target has {} entries",
target_data.len()
)));
}
let weight_data = match weight {
Some(weights) => Some(weights.to_vec()?),
None => None,
};
let divisor = num_classes.saturating_sub(1).max(1) as f32;
let mut selector = vec![0.0f32; batch_size * num_classes];
let mut off_class = vec![0.0f32; batch_size * num_classes];
let mut sample_scale = vec![0.0f32; batch_size];
for (b, &target_index) in target_data.iter().enumerate() {
let true_class = target_index as usize;
if true_class >= num_classes {
continue;
}
let row = b * num_classes;
if let Some(slot) = selector.get_mut(row + true_class) {
*slot = 1.0;
}
for c in 0..num_classes {
if c != true_class {
if let Some(slot) = off_class.get_mut(row + c) {
*slot = 1.0;
}
}
}
let class_weight = weight_data
.as_ref()
.and_then(|weights| weights.get(true_class).copied())
.unwrap_or(1.0);
if let Some(slot) = sample_scale.get_mut(b) {
*slot = class_weight / divisor;
}
}
let selector_tensor =
Tensor::from_data(selector, vec![batch_size, num_classes], input.device())?;
let off_class_tensor =
Tensor::from_data(off_class, vec![batch_size, num_classes], input.device())?;
let scale_tensor = Tensor::from_data(sample_scale, vec![batch_size], input.device())?;
let true_score = input.mul_op(&selector_tensor)?.sum_dim(&[-1], true)?;
let hinge = input.sub(&true_score)?.add_scalar(margin)?.clamp_min(0.0)?;
let powered = if p == 1 { hinge } else { hinge.pow(p as f32)? };
let per_sample = powered
.mul_op(&off_class_tensor)?
.sum_dim(&[-1], false)?
.mul_op(&scale_tensor)?;
apply_reduction(&per_sample, reduction, None, &[])
}
pub fn multilabel_margin_loss(input: &Tensor, target: &Tensor, reduction: &str) -> Result<Tensor> {
let ones = torsh_tensor::creation::ones_like(input)?;
let margin_tensor = torsh_tensor::creation::full_like(input, 1.0)?;
let target_scores = input.mul_op(target)?;
let non_target_scores = input.mul_op(&ones.sub(target)?)?;
let margin = margin_tensor.sub(&target_scores)?.add(&non_target_scores)?;
let zeros = torsh_tensor::creation::zeros_like(input)?;
let loss = margin.maximum(&zeros)?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn mse_loss(input: &Tensor, target: &Tensor, reduction: &str) -> Result<Tensor> {
let diff = input.sub(target)?;
let squared_diff = diff.mul_op(&diff)?;
apply_reduction(&squared_diff, reduction, None, &[])
}
pub fn l1_loss(input: &Tensor, target: &Tensor, reduction: &str) -> Result<Tensor> {
let diff = input.sub(target)?;
let abs_diff = diff.abs()?;
apply_reduction(&abs_diff, reduction, None, &[])
}
fn banded_abs_error(input: &Tensor, target: &Tensor, threshold: f32) -> Result<(Tensor, Tensor)> {
let abs_diff = input.sub(target)?.abs()?;
let inside = abs_diff.clamp_max(threshold)?;
let outside = abs_diff.sub(&inside)?;
Ok((inside, outside))
}
pub fn smooth_l1_loss(
input: &Tensor,
target: &Tensor,
beta: f32,
reduction: &str,
) -> Result<Tensor> {
if beta == 0.0 {
let abs_diff = input.sub(target)?.abs()?;
return apply_reduction(&abs_diff, reduction, None, &[]);
}
let (inside, outside) = banded_abs_error(input, target, beta)?;
let quadratic = inside.square()?.mul_scalar(0.5)?.div_scalar(beta)?;
let loss = quadratic.add(&outside)?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn huber_loss(input: &Tensor, target: &Tensor, delta: f32, reduction: &str) -> Result<Tensor> {
let (inside, outside) = banded_abs_error(input, target, delta)?;
let quadratic = inside.square()?.mul_scalar(0.5)?;
let linear = outside.mul_scalar(delta)?;
let loss = quadratic.add(&linear)?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn kl_div(
input: &Tensor,
target: &Tensor,
reduction: &str,
log_target: bool,
) -> Result<Tensor> {
let eps = 1e-8f32;
let eps_tensor = torsh_tensor::creation::full_like(target, eps)?;
let (target_probs, log_target_probs) = if log_target {
let target_probs = target.exp()?;
(target_probs, target.clone())
} else {
let stable_target = target.add(&eps_tensor)?;
let log_target_probs = stable_target.log()?;
(target.clone(), log_target_probs)
};
let log_ratio = log_target_probs.sub(input)?;
let kl_elements = target_probs.mul_op(&log_ratio)?;
match reduction {
"mean" => {
kl_elements.mean(None, false)
}
"sum" => {
kl_elements.sum()
}
"batchmean" => {
let batch_size = input.shape().dims()[0] as f32;
let total_sum = kl_elements.sum()?;
let batch_size_tensor = torsh_tensor::creation::full(&[1], batch_size)?;
total_sum.div(&batch_size_tensor)
}
"none" => {
Ok(kl_elements)
}
_ => Err(TorshError::InvalidArgument(format!(
"Unknown reduction: {}. Expected 'mean', 'sum', 'batchmean', or 'none'",
reduction
))),
}
}
pub fn nll_loss(
input: &Tensor,
target: &Tensor<i64>,
weight: Option<&Tensor>,
ignore_index: Option<i64>,
reduction: &str,
) -> Result<Tensor> {
let input_shape_binding = input.shape();
let input_dims = input_shape_binding.dims();
let (batch_size, num_classes) = match input_dims {
[batch, classes] => (*batch, *classes),
_ => {
return Err(TorshError::InvalidShape(format!(
"nll_loss expects a 2-D input [N, C], got shape {input_dims:?}"
)))
}
};
let target_data = target.to_vec()?;
if target_data.len() < batch_size {
return Err(TorshError::InvalidShape(format!(
"nll_loss expects one target per sample: input has {batch_size} rows \
but target has {} entries",
target_data.len()
)));
}
let weight_data = match weight {
Some(weights) => Some(weights.to_vec()?),
None => None,
};
let mut selector = vec![0.0f32; batch_size * num_classes];
for (b, &target_index) in target_data.iter().take(batch_size).enumerate() {
if Some(target_index) == ignore_index {
continue;
}
let target_class = target_index as usize;
if target_class >= num_classes {
continue;
}
let class_weight = match weight_data.as_ref() {
Some(values) => values.get(target_class).copied().unwrap_or(1.0),
None => 1.0,
};
if let Some(slot) = selector.get_mut(b * num_classes + target_class) {
*slot = class_weight;
}
}
let selector_tensor =
Tensor::from_data(selector, vec![batch_size, num_classes], input.device())?;
let loss_tensor = input
.mul_op(&selector_tensor)?
.neg()?
.sum_dim(&[-1], false)?;
apply_reduction(&loss_tensor, reduction, ignore_index, &target_data)
}
pub fn focal_loss(
input: &Tensor,
target: &Tensor<i64>,
alpha: Option<f32>,
gamma: f32,
reduction: &str,
) -> Result<Tensor> {
let input_shape = input.shape();
let input_dims = input_shape.dims();
let (batch_size, num_classes) = match input_dims {
[batch, classes] => (*batch, *classes),
_ => {
return Err(torsh_core::error::TorshError::InvalidShape(format!(
"Input must be 2D [batch_size, num_classes], got shape {:?}",
input_dims
)))
}
};
let target_data = target.to_vec()?;
if target_data.len() < batch_size {
return Err(torsh_core::error::TorshError::InvalidShape(format!(
"focal_loss expects one target per sample: input has {batch_size} rows \
but target has {} entries",
target_data.len()
)));
}
let mut selector = vec![0.0f32; batch_size * num_classes];
for (b, &target_index) in target_data.iter().take(batch_size).enumerate() {
let target_class = target_index as usize;
if target_class >= num_classes {
return Err(torsh_core::error::TorshError::InvalidArgument(format!(
"Target class {} out of range for {} classes",
target_class, num_classes
)));
}
if let Some(slot) = selector.get_mut(b * num_classes + target_class) {
*slot = 1.0;
}
}
let selector_tensor =
Tensor::from_data(selector, vec![batch_size, num_classes], input.device())?;
let log_probs = input.log_softmax(-1)?;
let log_pt = log_probs.mul_op(&selector_tensor)?.sum_dim(&[-1], false)?;
let pt = log_pt.exp()?;
let alpha_weight = alpha.unwrap_or(1.0);
let focal_weight = pt
.neg()?
.add_scalar(1.0)?
.pow(gamma)?
.mul_scalar(alpha_weight)?;
let per_sample = focal_weight.mul_op(&log_pt)?.neg()?;
match reduction {
"mean" => per_sample.mean(None, false)?.view(&[1]),
"sum" => per_sample.sum()?.view(&[1]),
"none" => Ok(per_sample),
_ => Err(TorshError::InvalidArgument(format!(
"Invalid reduction mode: '{}'. Expected 'mean', 'sum', or 'none'",
reduction
))),
}
}
pub fn triplet_margin_loss(
anchor: &Tensor,
positive: &Tensor,
negative: &Tensor,
margin: f32,
p: f32,
reduction: &str,
) -> Result<Tensor> {
let anchor_shape_obj = anchor.shape();
let anchor_shape = anchor_shape_obj.dims();
if anchor_shape.is_empty() {
return Err(TorshError::InvalidShape(
"triplet_margin_loss expects operands with a leading batch axis, got a 0-D tensor"
.to_string(),
));
}
let feature_axes = feature_axes_of(anchor_shape);
let dist_ap = p_norm_distance(anchor, positive, p, &feature_axes)?;
let dist_an = p_norm_distance(anchor, negative, p, &feature_axes)?;
let per_sample = dist_ap.sub(&dist_an)?.add_scalar(margin)?.clamp_min(0.0)?;
match reduction {
"mean" => per_sample.mean(None, false)?.view(&[1]),
"sum" => per_sample.sum()?.view(&[1]),
"none" => Ok(per_sample),
_ => Err(TorshError::InvalidArgument(format!(
"Invalid reduction mode: {}. Expected 'mean', 'sum', or 'none'",
reduction
))),
}
}
pub fn contrastive_loss(
output1: &Tensor,
output2: &Tensor,
target: &Tensor,
margin: f32,
reduction: &str,
) -> Result<Tensor> {
let output1_shape_obj = output1.shape();
let output1_shape = output1_shape_obj.dims();
let Some(&batch_size) = output1_shape.first() else {
return Err(TorshError::InvalidShape(
"contrastive_loss expects embeddings with a leading batch axis, got a 0-D tensor"
.to_string(),
));
};
let feature_axes = feature_axes_of(output1_shape);
let squared = output1.sub(output2)?.square()?;
let dist_squared = reduce_features(squared, &feature_axes)?;
let dist = dist_squared.clamp_min(DISTANCE_FLOOR)?.sqrt()?;
let target_data = target.to_vec()?;
if target_data.len() != batch_size {
return Err(TorshError::InvalidShape(format!(
"contrastive_loss expects one label per sample: the embeddings have {batch_size} \
rows but target has {} entries",
target_data.len()
)));
}
let similar: Vec<f32> = target_data
.iter()
.map(|&label| if label > 0.5 { 1.0 } else { 0.0 })
.collect();
let dissimilar: Vec<f32> = similar.iter().map(|&value| 1.0 - value).collect();
let similar_mask = Tensor::from_data(similar, vec![batch_size], output1.device())?;
let dissimilar_mask = Tensor::from_data(dissimilar, vec![batch_size], output1.device())?;
let repulsion = dist.neg()?.add_scalar(margin)?.clamp_min(0.0)?.square()?;
let per_sample = dist_squared
.mul_op(&similar_mask)?
.add(&repulsion.mul_op(&dissimilar_mask)?)?;
match reduction {
"mean" => per_sample.mean(None, false)?.view(&[1]),
"sum" => per_sample.sum()?.view(&[1]),
"none" => Ok(per_sample),
_ => Err(TorshError::InvalidArgument(format!(
"Invalid reduction mode: {}. Expected 'mean', 'sum', or 'none'",
reduction
))),
}
}
pub fn cosine_embedding_loss(
input1: &Tensor,
input2: &Tensor,
target: &Tensor,
margin: f32,
reduction: &str,
) -> Result<Tensor> {
let input1_shape_obj = input1.shape();
let input1_shape = input1_shape_obj.dims();
if input1_shape.is_empty() {
return Err(TorshError::InvalidShape(
"cosine_embedding_loss expects embeddings with a feature axis, got a 0-D tensor"
.to_string(),
));
}
let feature_axis = [input1_shape.len() as i32 - 1];
let dot_product = input1.mul_op(input2)?.sum_dim(&feature_axis, false)?;
let norm1_squared = input1.square()?.sum_dim(&feature_axis, false)?;
let norm2_squared = input2.square()?.sum_dim(&feature_axis, false)?;
let denominator = norm1_squared
.mul_op(&norm2_squared)?
.clamp_min(DISTANCE_FLOOR)?
.sqrt()?;
let cosine_sim = dot_product.div(&denominator)?;
let cosine_shape_obj = cosine_sim.shape();
let cosine_dims = cosine_shape_obj.dims().to_vec();
let sample_count: usize = cosine_dims.iter().product();
let target_data = target.to_vec()?;
if target_data.len() != sample_count {
return Err(TorshError::InvalidShape(format!(
"cosine_embedding_loss expects one label per sample: the embeddings yield \
{sample_count} cosine values but target has {} entries",
target_data.len()
)));
}
let positive: Vec<f32> = target_data
.iter()
.map(|&label| if label > 0.0 { 1.0 } else { 0.0 })
.collect();
let negative: Vec<f32> = positive.iter().map(|&value| 1.0 - value).collect();
let positive_mask = Tensor::from_data(positive, cosine_dims.clone(), input1.device())?;
let negative_mask = Tensor::from_data(negative, cosine_dims, input1.device())?;
let attraction = cosine_sim.neg()?.add_scalar(1.0)?;
let repulsion = cosine_sim.sub_scalar(margin)?.clamp_min(0.0)?;
let loss = attraction
.mul_op(&positive_mask)?
.add(&repulsion.mul_op(&negative_mask)?)?;
apply_reduction(&loss, reduction, None, &[])
}
const DISTANCE_FLOOR: f32 = 1e-12;
fn feature_axes_of(dims: &[usize]) -> Vec<i32> {
(1..dims.len() as i32).collect()
}
fn reduce_features(values: Tensor, feature_axes: &[i32]) -> Result<Tensor> {
if feature_axes.is_empty() {
Ok(values)
} else {
values.sum_dim(feature_axes, false)
}
}
fn p_norm_distance(x1: &Tensor, x2: &Tensor, p: f32, feature_axes: &[i32]) -> Result<Tensor> {
let powered = x1.sub(x2)?.abs()?.pow(p)?;
let summed = reduce_features(powered, feature_axes)?;
summed.clamp_min(DISTANCE_FLOOR)?.pow(1.0 / p)
}
fn apply_reduction(
loss: &Tensor,
reduction: &str,
ignore_index: Option<i64>,
target_data: &[i64],
) -> Result<Tensor> {
match reduction {
"mean" => {
if let Some(ignore_idx) = ignore_index {
let valid_count =
target_data.iter().filter(|&&idx| idx != ignore_idx).count() as f32;
if valid_count > 0.0 {
let sum = loss.sum()?;
let count_tensor = torsh_tensor::creation::full(&[1], valid_count)?;
sum.div(&count_tensor)
} else {
loss.mean(None, false)
}
} else {
loss.mean(None, false)
}
}
"sum" => loss.sum(),
"none" => Ok(loss.clone()),
_ => Err(TorshError::ComputeError(format!(
"Unknown reduction: {}",
reduction
))),
}
}
#[allow(dead_code)]
fn gather_target_probs(probs: &Tensor, target: &Tensor) -> Result<Tensor> {
let target_shape = target.shape();
let probs_shape = probs.shape();
if target_shape.dims().len() + 1 != probs_shape.dims().len() {
return Err(TorshError::InvalidArgument(
"Target and input tensor shapes are incompatible for gathering".to_string(),
));
}
let batch_size = target_shape.dims()[0];
let flat_size = target_shape.numel();
let mut result_data = Vec::with_capacity(flat_size);
let target_data = target.to_vec()?;
let probs_data = probs.to_vec()?;
let num_classes = probs_shape.dims()[probs_shape.dims().len() - 1];
for (i, &target_class) in target_data.iter().enumerate() {
let target_idx = target_class as usize;
if target_idx < num_classes {
let prob_idx = (i / batch_size) * num_classes + target_idx;
result_data.push(probs_data[prob_idx]);
} else {
result_data.push(0.0);
}
}
Tensor::from_vec(result_data, target_shape.dims())
}
#[allow(dead_code)]
fn pairwise_distance(x1: &Tensor, x2: &Tensor, p: f32) -> Result<Tensor> {
let diff = x1.sub(x2)?;
let abs_diff = diff.abs()?;
if p == 2.0 {
let squared = abs_diff.pow(2.0)?;
let sum_squared = squared.sum()?;
sum_squared.sqrt()
} else if p == 1.0 {
abs_diff.sum()
} else {
let powered = abs_diff.pow(p)?;
let sum_powered = powered.sum()?;
let inv_p = 1.0 / p;
sum_powered.pow(inv_p)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_triplet_margin_loss_basic() -> Result<()> {
let anchor = Tensor::from_vec(vec![1.0, 1.0], &[1, 2])?;
let positive = Tensor::from_vec(vec![1.1, 1.1], &[1, 2])?;
let negative = Tensor::from_vec(vec![5.0, 5.0], &[1, 2])?;
let loss = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "mean")?;
let loss_data = loss.to_vec()?;
assert!(
loss_data[0] < 0.1,
"Loss should be near zero when constraint is satisfied"
);
Ok(())
}
#[test]
fn test_triplet_margin_loss_violation() -> Result<()> {
let anchor = Tensor::from_vec(vec![0.0, 0.0], &[1, 2])?;
let positive = Tensor::from_vec(vec![2.0, 0.0], &[1, 2])?;
let negative = Tensor::from_vec(vec![1.0, 0.0], &[1, 2])?;
let loss = triplet_margin_loss(&anchor, &positive, &negative, 0.5, 2.0, "mean")?;
let loss_data = loss.to_vec()?;
assert!((loss_data[0] - 1.5).abs() < 1e-5, "Loss should be 1.5");
Ok(())
}
#[test]
fn test_triplet_margin_loss_batch() -> Result<()> {
let anchor = Tensor::from_vec(
vec![
0.0, 0.0, 1.0, 1.0, ],
&[2, 2],
)?;
let positive = Tensor::from_vec(
vec![
0.1, 0.1, 1.1, 1.1, ],
&[2, 2],
)?;
let negative = Tensor::from_vec(
vec![
5.0, 5.0, 6.0, 6.0, ],
&[2, 2],
)?;
let loss = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "none")?;
assert_eq!(loss.shape().dims(), &[2]);
let loss_mean = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "mean")?;
assert_eq!(loss_mean.shape().dims(), &[1]);
Ok(())
}
#[test]
fn test_contrastive_loss_similar_pairs() -> Result<()> {
let output1 = Tensor::from_vec(vec![1.0, 2.0], &[1, 2])?;
let output2 = Tensor::from_vec(vec![1.1, 2.1], &[1, 2])?;
let target = Tensor::from_vec(vec![1.0], &[1])?;
let loss = contrastive_loss(&output1, &output2, &target, 2.0, "mean")?;
let loss_data = loss.to_vec()?;
assert!(
(loss_data[0] - 0.02).abs() < 1e-5,
"Loss for similar pair should be distance squared"
);
Ok(())
}
#[test]
fn test_contrastive_loss_dissimilar_pairs() -> Result<()> {
let output1 = Tensor::from_vec(vec![0.0, 0.0], &[1, 2])?;
let output2 = Tensor::from_vec(vec![0.5, 0.0], &[1, 2])?;
let target = Tensor::from_vec(vec![0.0], &[1])?;
let loss = contrastive_loss(&output1, &output2, &target, 2.0, "mean")?;
let loss_data = loss.to_vec()?;
assert!(
(loss_data[0] - 2.25).abs() < 1e-5,
"Loss for dissimilar pair should be (margin - dist)^2"
);
Ok(())
}
#[test]
fn test_contrastive_loss_dissimilar_beyond_margin() -> Result<()> {
let output1 = Tensor::from_vec(vec![0.0, 0.0], &[1, 2])?;
let output2 = Tensor::from_vec(vec![5.0, 0.0], &[1, 2])?;
let target = Tensor::from_vec(vec![0.0], &[1])?;
let loss = contrastive_loss(&output1, &output2, &target, 2.0, "mean")?;
let loss_data = loss.to_vec()?;
assert!(
loss_data[0] < 1e-5,
"Loss should be zero when dissimilar pairs are beyond margin"
);
Ok(())
}
#[test]
fn test_contrastive_loss_batch() -> Result<()> {
let output1 = Tensor::from_vec(
vec![
0.0, 0.0, 1.0, 1.0, ],
&[2, 2],
)?;
let output2 = Tensor::from_vec(
vec![
0.1, 0.0, 1.0, 2.0, ],
&[2, 2],
)?;
let target = Tensor::from_vec(
vec![
1.0, 0.0, ],
&[2],
)?;
let loss = contrastive_loss(&output1, &output2, &target, 2.0, "none")?;
assert_eq!(loss.shape().dims(), &[2]);
let loss_mean = contrastive_loss(&output1, &output2, &target, 2.0, "mean")?;
assert_eq!(loss_mean.shape().dims(), &[1]);
Ok(())
}
#[test]
fn test_reduction_modes() -> Result<()> {
let anchor = Tensor::from_vec(vec![0.0, 0.0, 1.0, 1.0], &[2, 2])?;
let positive = Tensor::from_vec(vec![0.1, 0.1, 1.1, 1.1], &[2, 2])?;
let negative = Tensor::from_vec(vec![5.0, 5.0, 6.0, 6.0], &[2, 2])?;
let loss_none = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "none")?;
assert_eq!(
loss_none.shape().dims(),
&[2],
"none reduction should return batch_size losses"
);
let loss_mean = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "mean")?;
assert_eq!(
loss_mean.shape().dims(),
&[1],
"mean reduction should return scalar"
);
let loss_sum = triplet_margin_loss(&anchor, &positive, &negative, 1.0, 2.0, "sum")?;
assert_eq!(
loss_sum.shape().dims(),
&[1],
"sum reduction should return scalar"
);
Ok(())
}
}
pub fn dice_loss(input: &Tensor, target: &Tensor, smooth: f32, reduction: &str) -> Result<Tensor> {
if input.shape().dims() != target.shape().dims() {
return Err(TorshError::ShapeMismatch {
expected: target.shape().dims().to_vec(),
got: input.shape().dims().to_vec(),
});
}
let intersection_sum = input.mul_op(target)?.sum()?;
let input_sum = input.sum()?;
let target_sum = target.sum()?;
let numerator = intersection_sum.mul_scalar(2.0)?.add_scalar(smooth)?;
let denominator = input_sum.add(&target_sum)?.add_scalar(smooth)?;
let dice_coeff = numerator.div(&denominator)?;
let loss = dice_coeff.neg()?.add_scalar(1.0)?.view(&[1])?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn tversky_loss(
input: &Tensor,
target: &Tensor,
alpha: f32,
beta: f32,
smooth: f32,
reduction: &str,
) -> Result<Tensor> {
if input.shape().dims() != target.shape().dims() {
return Err(TorshError::ShapeMismatch {
expected: target.shape().dims().to_vec(),
got: input.shape().dims().to_vec(),
});
}
if alpha + beta > 1.0 {
return Err(TorshError::InvalidArgument(
"alpha + beta should be <= 1.0 for Tversky loss".to_string(),
));
}
let tp = input.mul_op(target)?.sum()?;
let ones = torsh_tensor::creation::ones_like(target)?;
let fp = input.mul_op(&ones.sub(target)?)?.sum()?;
let fn_tensor = ones.sub(input)?.mul_op(target)?.sum()?;
let numerator = tp.add_scalar(smooth)?;
let denominator = tp
.add(&fp.mul_scalar(alpha)?)?
.add(&fn_tensor.mul_scalar(beta)?)?
.add_scalar(smooth)?;
let tversky_index = numerator.div(&denominator)?;
let loss = tversky_index.neg()?.add_scalar(1.0)?.view(&[1])?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn wing_loss(
input: &Tensor,
target: &Tensor,
width: f32,
curvature: f32,
reduction: &str,
) -> Result<Tensor> {
if input.shape().dims() != target.shape().dims() {
return Err(TorshError::ShapeMismatch {
expected: target.shape().dims().to_vec(),
got: input.shape().dims().to_vec(),
});
}
let (inside, outside) = banded_abs_error(input, target, width)?;
let logarithmic = inside
.div_scalar(curvature)?
.add_scalar(1.0)?
.log()?
.mul_scalar(width)?;
let loss = logarithmic.add(&outside)?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn center_loss(
features: &Tensor,
labels: &Tensor<i64>,
centers: &Tensor,
reduction: &str,
) -> Result<Tensor> {
let features_shape_binding = features.shape();
let features_shape = features_shape_binding.dims();
let (batch_size, feature_dim) = match features_shape {
[batch, dim] => (*batch, *dim),
_ => {
return Err(TorshError::InvalidShape(format!(
"center_loss expects 2-D features [N, D], got shape {features_shape:?}"
)))
}
};
let centers_shape_binding = centers.shape();
let centers_shape = centers_shape_binding.dims();
let num_classes = centers_shape[0];
if centers_shape.len() != 2 || centers_shape[1] != feature_dim {
return Err(TorshError::ShapeMismatch {
expected: vec![num_classes, feature_dim],
got: centers_shape.to_vec(),
});
}
let labels_vec: Vec<i64> = labels.to_vec()?;
if labels_vec.len() != batch_size {
return Err(TorshError::InvalidShape(format!(
"center_loss expects one label per sample: features have {batch_size} rows \
but labels have {} entries",
labels_vec.len()
)));
}
let mut assignment = vec![0.0f32; batch_size * num_classes];
for (b, &label) in labels_vec.iter().enumerate() {
let label_idx = label as usize;
if label_idx >= num_classes {
return Err(TorshError::InvalidArgument(format!(
"Label {} out of range for {} classes",
label, num_classes
)));
}
if let Some(slot) = assignment.get_mut(b * num_classes + label_idx) {
*slot = 1.0;
}
}
let assignment_tensor =
Tensor::from_data(assignment, vec![batch_size, num_classes], features.device())?;
let selected = assignment_tensor.matmul(centers)?;
let diff = features.sub(&selected)?;
let loss = diff.square()?.sum_dim(&[-1], false)?.mul_scalar(0.5)?;
apply_reduction(&loss, reduction, None, &[])
}
pub fn infonce_loss(
anchor: &Tensor,
positive: &Tensor,
negatives: &Tensor,
temperature: f32,
reduction: &str,
) -> Result<Tensor> {
let anchor_shape_binding = anchor.shape();
let anchor_shape = anchor_shape_binding.dims();
let (batch_size, embedding_dim) = match anchor_shape {
[batch, dim] => (*batch, *dim),
_ => {
return Err(TorshError::InvalidShape(format!(
"infonce_loss expects a 2-D anchor [N, D], got shape {anchor_shape:?}"
)))
}
};
let positive_shape = positive.shape();
if positive_shape.dims() != anchor_shape {
return Err(TorshError::ShapeMismatch {
expected: anchor_shape.to_vec(),
got: positive_shape.dims().to_vec(),
});
}
let negatives_shape_binding = negatives.shape();
let negatives_shape = negatives_shape_binding.dims();
if negatives_shape.len() != 2 || negatives_shape[1] != embedding_dim {
return Err(TorshError::ShapeMismatch {
expected: vec![negatives_shape[0], embedding_dim],
got: negatives_shape.to_vec(),
});
}
let num_negatives = negatives_shape[0];
let anchor_norm = row_norms(anchor)?;
let positive_norm = row_norms(positive)?;
let positive_similarity = anchor
.mul_op(positive)?
.sum_dim(&[-1], true)?
.div(&anchor_norm.mul_op(&positive_norm)?)?
.div_scalar(temperature)?;
let logits = if num_negatives == 0 {
positive_similarity
} else {
let negative_norm = row_norms(negatives)?;
let negative_similarity = anchor
.matmul(&negatives.transpose(0, 1)?)?
.div(&anchor_norm.matmul(&negative_norm.transpose(0, 1)?)?)?
.div_scalar(temperature)?;
Tensor::cat(&[&positive_similarity, &negative_similarity], 1)?
};
let log_probs = logits.log_softmax(-1)?;
let row_width = num_negatives + 1;
let mut selector = vec![0.0f32; batch_size * row_width];
for b in 0..batch_size {
if let Some(slot) = selector.get_mut(b * row_width) {
*slot = 1.0;
}
}
let selector_tensor =
Tensor::from_data(selector, vec![batch_size, row_width], anchor.device())?;
let loss = log_probs
.mul_op(&selector_tensor)?
.sum_dim(&[-1], false)?
.neg()?;
apply_reduction(&loss, reduction, None, &[])
}
fn row_norms(rows: &Tensor) -> Result<Tensor> {
let norm = rows.square()?.sum_dim(&[-1], true)?.sqrt()?;
let norm_values = norm.to_vec()?;
if !norm_values.iter().any(|value| *value == 0.0) {
return Ok(norm);
}
let fallback: Vec<f32> = norm_values
.iter()
.map(|value| if *value == 0.0 { 1.0 } else { 0.0 })
.collect();
let fallback_tensor = Tensor::from_data(fallback, norm.shape().dims().to_vec(), rows.device())?;
norm.add(&fallback_tensor)
}