use burn_std::FloatDType;
use crate::{AsIndex, check::unwrap_dim_index, tensor::Tensor};
use super::vector_norm::l2_norm_impl;
pub fn cosine_similarity<const D: usize>(
x1: Tensor<D>,
x2: Tensor<D>,
dim: impl AsIndex,
eps: Option<f64>,
) -> Tensor<D> {
let dim = unwrap_dim_index(dim.try_dim_index(D), "Cosine Similarity");
let eps = eps.unwrap_or_else(|| {
x1.dtype()
.finfo()
.unwrap_or(FloatDType::F32.finfo())
.min_positive
});
let dot_product = (x1.clone() * x2.clone()).sum_dim(dim);
let norm_x1 = l2_norm_impl(x1, dim);
let norm_x2 = l2_norm_impl(x2, dim);
let denominator = norm_x1.clamp_min(eps) * norm_x2.clamp_min(eps);
dot_product / denominator
}