kopitiam_tensor/tensor/
softmax.rs1use kopitiam_core::{DType, Error, Result};
4
5use crate::storage::Storage;
6
7use super::Tensor;
8
9impl Tensor {
10 pub fn softmax(&self, axis: usize) -> Result<Tensor> {
26 self.require_dtype(DType::F32)?;
27 let rank = self.rank();
28 if axis >= rank {
29 return Err(Error::IndexOutOfBounds { dim: axis, index: axis, len: rank });
30 }
31 let Storage::F32(data) = self.storage.as_ref() else { unreachable!() };
32 let contiguous: Vec<f32> = self.logical_offsets().map(|i| data[i]).collect();
33
34 let dims = self.shape.dims();
35 let axis_len = dims[axis];
36 let outer: usize = dims[..axis].iter().product();
37 let inner: usize = dims[axis + 1..].iter().product();
38
39 let mut out = vec![0f32; contiguous.len()];
40 for o in 0..outer {
41 for inn in 0..inner {
42 let base = o * axis_len * inner + inn;
43 let at = |a: usize| contiguous[base + a * inner];
44
45 let max = (0..axis_len).map(at).fold(f32::NEG_INFINITY, f32::max);
46 let mut sum = 0f32;
47 for a in 0..axis_len {
48 let exp = (at(a) - max).exp();
49 out[base + a * inner] = exp;
50 sum += exp;
51 }
52 for a in 0..axis_len {
53 out[base + a * inner] /= sum;
54 }
55 }
56 }
57 Tensor::from_f32(out, self.shape.clone())
58 }
59}
60
61#[cfg(test)]
62mod tests {
63 use super::*;
64
65 fn assert_close(a: f32, b: f32) {
66 assert!((a - b).abs() < 1e-5, "expected {b}, got {a}");
67 }
68
69 #[test]
70 fn softmax_rows_sum_to_one() {
71 let t = Tensor::from_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]).unwrap();
72 let out = t.softmax(1).unwrap().to_vec_f32().unwrap();
73 assert_close(out[0] + out[1] + out[2], 1.0);
74 assert_close(out[3] + out[4] + out[5], 1.0);
75 }
76
77 #[test]
78 fn softmax_matches_hand_computation_for_a_simple_row() {
79 let t = Tensor::from_f32(vec![0.0, 0.0], [2]).unwrap();
81 let out = t.softmax(0).unwrap().to_vec_f32().unwrap();
82 assert_close(out[0], 0.5);
83 assert_close(out[1], 0.5);
84 }
85
86 #[test]
87 fn softmax_preserves_relative_order() {
88 let t = Tensor::from_f32(vec![1.0, 3.0, 2.0], [3]).unwrap();
89 let out = t.softmax(0).unwrap().to_vec_f32().unwrap();
90 assert!(out[1] > out[2]);
91 assert!(out[2] > out[0]);
92 }
93
94 #[test]
99 fn large_magnitude_input_does_not_overflow_to_nan() {
100 let t = Tensor::from_f32(vec![1000.0, 1001.0], [2]).unwrap();
101 let out = t.softmax(0).unwrap().to_vec_f32().unwrap();
102 assert!(!out[0].is_nan() && !out[1].is_nan(), "softmax produced NaN: {out:?}");
103 assert_close(out[0] + out[1], 1.0);
104 assert_close(out[0], 0.268_941_4);
106 }
107
108 #[test]
109 fn softmax_operates_along_an_arbitrary_axis_of_a_3d_tensor() {
110 let t = Tensor::from_f32(vec![1.0, 2.0, 3.0, 4.0], [1, 2, 2]).unwrap();
112 let out = t.softmax(1).unwrap().to_vec_f32().unwrap();
113 assert_close(out[0] + out[2], 1.0);
115 assert_close(out[1] + out[3], 1.0);
116 }
117}