use trueno::Vector;
fn apply_vector_op<E>(
x: &[f32],
op: impl FnOnce(&Vector<f32>) -> Result<Vector<f32>, E>,
) -> Vec<f32> {
if x.is_empty() {
return vec![];
}
let vx = Vector::from_slice(x);
op(&vx).map_or_else(|_| vec![0.0; x.len()], |v| v.as_slice().to_vec())
}
#[must_use]
pub fn softmax(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::softmax)
}
#[must_use]
pub fn log_softmax(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::log_softmax)
}
#[must_use]
pub fn gelu(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::gelu)
}
#[must_use]
pub fn relu(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::relu)
}
#[must_use]
pub fn sigmoid(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::sigmoid)
}
#[must_use]
pub fn tanh_activation(x: &[f32]) -> Vec<f32> {
apply_vector_op(x, Vector::tanh)
}
#[cfg(test)]
mod tests {
use super::*;
const EPSILON: f32 = 1e-4;
fn approx_eq(a: f32, b: f32) -> bool {
(a - b).abs() < EPSILON
}
fn vec_approx_eq(a: &[f32], b: &[f32]) -> bool {
a.len() == b.len() && a.iter().zip(b).all(|(x, y)| approx_eq(*x, *y))
}
#[test]
fn test_softmax() {
let x = vec![1.0, 2.0, 3.0];
let result = softmax(&x);
assert_eq!(result.len(), 3);
let total: f32 = result.iter().sum();
assert!(approx_eq(total, 1.0));
assert!(result[0] < result[1]);
assert!(result[1] < result[2]);
}
#[test]
fn test_softmax_numerical_stability() {
let x = vec![1000.0, 1001.0, 1002.0];
let result = softmax(&x);
let total: f32 = result.iter().sum();
assert!(approx_eq(total, 1.0));
assert!(result.iter().all(|&v| v.is_finite()));
}
#[test]
fn test_softmax_empty() {
let x: Vec<f32> = vec![];
let result = softmax(&x);
assert!(result.is_empty());
}
#[test]
fn test_log_softmax() {
let x = vec![1.0, 2.0, 3.0];
let result = log_softmax(&x);
let softmax_result = softmax(&x);
let exp_log_softmax: Vec<f32> = result.iter().map(|v| v.exp()).collect();
assert!(vec_approx_eq(&exp_log_softmax, &softmax_result));
}
#[test]
fn test_gelu() {
let x = vec![-1.0, 0.0, 1.0];
let result = gelu(&x);
assert_eq!(result.len(), 3);
assert!(approx_eq(result[1], 0.0));
assert!(result[2] > 0.0);
}
#[test]
fn test_relu() {
let x = vec![-1.0, 0.0, 1.0, 2.0];
let result = relu(&x);
assert!(vec_approx_eq(&result, &[0.0, 0.0, 1.0, 2.0]));
}
#[test]
fn test_sigmoid() {
let x = vec![-100.0, 0.0, 100.0];
let result = sigmoid(&x);
assert!(result[0] < 0.01);
assert!(approx_eq(result[1], 0.5));
assert!(result[2] > 0.99);
}
#[test]
fn test_tanh() {
let x = vec![-100.0, 0.0, 100.0];
let result = tanh_activation(&x);
assert!(result[0] < -0.99);
assert!(approx_eq(result[1], 0.0));
assert!(result[2] > 0.99);
}
#[test]
fn test_log_softmax_empty() {
let x: Vec<f32> = vec![];
let result = log_softmax(&x);
assert!(result.is_empty());
}
#[test]
fn test_gelu_empty() {
let x: Vec<f32> = vec![];
let result = gelu(&x);
assert!(result.is_empty());
}
#[test]
fn test_relu_empty() {
let x: Vec<f32> = vec![];
let result = relu(&x);
assert!(result.is_empty());
}
#[test]
fn test_sigmoid_empty() {
let x: Vec<f32> = vec![];
let result = sigmoid(&x);
assert!(result.is_empty());
}
#[test]
fn test_tanh_empty() {
let x: Vec<f32> = vec![];
let result = tanh_activation(&x);
assert!(result.is_empty());
}
#[test]
fn test_softmax_single() {
let x = vec![1.0];
let result = softmax(&x);
assert_eq!(result.len(), 1);
assert!(approx_eq(result[0], 1.0)); }
#[test]
fn test_log_softmax_single() {
let x = vec![1.0];
let result = log_softmax(&x);
assert_eq!(result.len(), 1);
assert!(approx_eq(result[0], 0.0)); }
#[test]
fn test_gelu_positive() {
let x = vec![0.5, 1.0, 2.0, 3.0];
let result = gelu(&x);
assert!(result[3] > 2.9);
assert!(result.iter().all(|&v| v > 0.0));
}
#[test]
fn test_gelu_negative() {
let x = vec![-3.0, -2.0, -1.0, -0.5];
let result = gelu(&x);
assert!(result.iter().all(|&v| v > -0.5));
}
#[test]
fn test_relu_all_positive() {
let x = vec![1.0, 2.0, 3.0, 4.0];
let result = relu(&x);
assert!(vec_approx_eq(&result, &x));
}
#[test]
fn test_relu_all_negative() {
let x = vec![-1.0, -2.0, -3.0, -4.0];
let result = relu(&x);
assert!(vec_approx_eq(&result, &[0.0, 0.0, 0.0, 0.0]));
}
#[test]
fn test_sigmoid_gradient_region() {
let x = vec![-2.0, -1.0, 0.0, 1.0, 2.0];
let result = sigmoid(&x);
for i in 1..result.len() {
assert!(result[i] > result[i - 1]);
}
}
#[test]
fn test_tanh_symmetry() {
let x = vec![-2.0, -1.0, 1.0, 2.0];
let result = tanh_activation(&x);
assert!(approx_eq(result[0], -result[3]));
assert!(approx_eq(result[1], -result[2]));
}
#[test]
fn test_softmax_uniform() {
let x = vec![1.0, 1.0, 1.0, 1.0];
let result = softmax(&x);
for &v in &result {
assert!(approx_eq(v, 0.25));
}
}
#[test]
fn test_log_softmax_numerical_stability() {
let x = vec![1000.0, 1001.0, 1002.0];
let result = log_softmax(&x);
assert!(result.iter().all(|&v| v.is_finite()));
assert!(result.iter().all(|&v| v <= 0.0));
}
}